mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 23:56:36 +08:00
Compare commits
364 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 841d43344a | |||
| 5efe790937 | |||
| eec6cb4298 | |||
| 352fc82907 | |||
| 87722e461c | |||
| 701b4011cb | |||
| d128d2f657 | |||
| 400a40fe80 | |||
| 103290ed35 | |||
| 363e714603 | |||
| afd1258fcd | |||
| e69082a596 | |||
| d30363d164 | |||
| 8e1a87bf5a | |||
| 6180b5a198 | |||
| 61d95ab5d5 | |||
| c27be19915 | |||
| f62a35c3f9 | |||
| 2a1caf32c4 | |||
| fdcc30a493 | |||
| efaffb0475 | |||
| 4954526cbc | |||
| 9d50071915 | |||
| ceceee6ebd | |||
| 11051f5517 | |||
| ff2c7c4959 | |||
| 6364b96935 | |||
| 409f0a232a | |||
| f79994e0e0 | |||
| 9fdb16d035 | |||
| 53b632a6f7 | |||
| bf7b2a0740 | |||
| 8475bc27bb | |||
| 5d01572eff | |||
| a353faaa71 | |||
| e7b25004ba | |||
| 3826cb02c0 | |||
| a1fee8e432 | |||
| aafdb78482 | |||
| 16b545d8cd | |||
| 45065178b8 | |||
| 8ebde9dca9 | |||
| 0a1ec60750 | |||
| 9ec35d2f2f | |||
| c914040b7d | |||
| 822362c44c | |||
| 80f5935b76 | |||
| 433c8aab13 | |||
| 949dfcd42d | |||
| 1580e4ee10 | |||
| 32e4f0f514 | |||
| ca3a643ef7 | |||
| ce2b234843 | |||
| ac3506847c | |||
| 6e3d604618 | |||
| 4417ece7cd | |||
| bab4371ba7 | |||
| 960c97cee4 | |||
| 9d05d75fd6 | |||
| c137bdcc63 | |||
| b7065f6e99 | |||
| bd4e1f66cb | |||
| 9f0670f4d0 | |||
| 32e338d295 | |||
| f6a753baa3 | |||
| 322a10bb9d | |||
| 27c13d6c47 | |||
| f45f96063a | |||
| 1780be73b9 | |||
| addf8e2089 | |||
| 75cd60ea3e | |||
| fe42a77409 | |||
| ce9abf457f | |||
| 3c57a5ac84 | |||
| ff7c91d277 | |||
| 1498f3052d | |||
| 6b1264ae90 | |||
| 18445ec063 | |||
| 08bc91e5c9 | |||
| 2f424bea31 | |||
| bc75ed745d | |||
| 78fb9a31d6 | |||
| d3ed2e8856 | |||
| db21ce6bb4 | |||
| 76ad841231 | |||
| d0535707dc | |||
| 6458b5af00 | |||
| 555039e028 | |||
| 7134253b2c | |||
| 58d29b440a | |||
| 6a3a9add08 | |||
| 6d986524f1 | |||
| 92a8fed796 | |||
| b314192621 | |||
| 1e5f9bfb04 | |||
| 5972378897 | |||
| 455900ba41 | |||
| 85e57213ee | |||
| 1377061234 | |||
| ea21a7deef | |||
| 9b98194a0a | |||
| 2df061a19f | |||
| 46a60376c4 | |||
| 9de240f034 | |||
| 41ef814643 | |||
| 6c7b4817f9 | |||
| 7507507fd9 | |||
| ff94406945 | |||
| 9aa13c4dfb | |||
| 4a8c400944 | |||
| 18e7ec94a8 | |||
| 02f2a1c8b3 | |||
| 681a0bef48 | |||
| 67bf5be0f2 | |||
| efb613b0b5 | |||
| 8124e59de5 | |||
| ac8c293ff3 | |||
| 4e38b73cac | |||
| 8f336377f6 | |||
| 5f78dd66fc | |||
| f2ee939006 | |||
| 23d2060742 | |||
| bb0da0b769 | |||
| e51af4be1f | |||
| a82f3a75b0 | |||
| 7d07fe08b7 | |||
| 375877b223 | |||
| 004daeadb6 | |||
| 3bcb80d7a2 | |||
| 5ff9621227 | |||
| 84db9711bc | |||
| 2f97e892d5 | |||
| 17fd1e4ad4 | |||
| e56dd898ef | |||
| 06bb8b3b04 | |||
| 42a775c3bb | |||
| 05c3b5842e | |||
| 149e10ee66 | |||
| e194813f3b | |||
| f1cad30f44 | |||
| 2e05df288b | |||
| 3e5bb8fc0b | |||
| d1e3c59537 | |||
| 8b8ebb6092 | |||
| 0195a2a01b | |||
| ad9b336fb9 | |||
| 30d9552207 | |||
| 5e96a8de72 | |||
| 69faeaa9a6 | |||
| e8bfe52104 | |||
| 9767cc3247 | |||
| e8a7f999c8 | |||
| 4d4f5f8b1f | |||
| d2a425d761 | |||
| 673d38a089 | |||
| 2e8c0530a9 | |||
| 32ee511eac | |||
| f410640862 | |||
| 6427b830ea | |||
| 5e7bf3ba5c | |||
| 27d6691232 | |||
| cc4b8a916a | |||
| d9dd5131b2 | |||
| a98c9f4f59 | |||
| 5cb935e0e5 | |||
| 6f59e4be0c | |||
| 647446a2a2 | |||
| 42d6249af5 | |||
| de9ab51def | |||
| a2ec08f033 | |||
| f8809d73fb | |||
| 413081f72a | |||
| e5339a8072 | |||
| fbb4d82a44 | |||
| 508a37a84c | |||
| d60655045a | |||
| 31ef861504 | |||
| f1bdb2e2ef | |||
| 61b71a11c7 | |||
| 4c69ff491d | |||
| 0ad4904e20 | |||
| bd30b61018 | |||
| e0dd70a054 | |||
| 4966a8aad1 | |||
| 3e11549370 | |||
| addf83a249 | |||
| c3e35fd416 | |||
| 775dfe19f1 | |||
| db3b2f651b | |||
| 669323f926 | |||
| 7202b69e4e | |||
| 31977a62e6 | |||
| 87479c2ac1 | |||
| ffda0fb71a | |||
| 9c0e7341c3 | |||
| 1db5452be9 | |||
| c10f894afd | |||
| 7fb75baa73 | |||
| 15e6cd69eb | |||
| f6eb88d75e | |||
| f45b580984 | |||
| 4f50c47550 | |||
| 2e1d75dc36 | |||
| f496f58a4d | |||
| 32474bec20 | |||
| 581cda7edc | |||
| 96aebb8d61 | |||
| 735fd40786 | |||
| a3b0bf4898 | |||
| 9703e4a081 | |||
| a43653f252 | |||
| 348900de01 | |||
| b93c259fac | |||
| 2e3d5c9249 | |||
| c8c1841058 | |||
| 1c596fae4b | |||
| 2ff52e3275 | |||
| 7efb49bdab | |||
| a00b20abf3 | |||
| 1450b25475 | |||
| b815be54b8 | |||
| 75edeb9afa | |||
| 7c54192055 | |||
| 7ba68778c1 | |||
| 7b736b2e60 | |||
| ef613c1518 | |||
| b62df6ffa3 | |||
| be9d8773ce | |||
| 1c10347357 | |||
| 5bd21e2ac1 | |||
| e38335973d | |||
| 95929bf82e | |||
| 9cf9f4f1f7 | |||
| ae8dbdd77f | |||
| 05bd6a686d | |||
| b8193417f5 | |||
| 15e4508be4 | |||
| 634c6cd620 | |||
| 4eaecb289b | |||
| 98a9e5c666 | |||
| d244920dd4 | |||
| 77e4387b35 | |||
| 7a40ddb1ef | |||
| d33814e18c | |||
| cf51b305b0 | |||
| 9ffeb83753 | |||
| 2f40cf29d4 | |||
| a92eb168aa | |||
| de21a55f37 | |||
| b01dbdb6e5 | |||
| a645cc699b | |||
| 528f912aac | |||
| 8bf30a157f | |||
| 58abba7fc0 | |||
| d8cd4b404c | |||
| 9e979aa82a | |||
| 5caaaf6092 | |||
| f23d1c2afd | |||
| 5e00cbf131 | |||
| 975948dcf6 | |||
| a9eac6d01f | |||
| 6e8406f439 | |||
| db3577afa9 | |||
| 7285717e34 | |||
| de6911f219 | |||
| e5ce0501a2 | |||
| 25a87e25c5 | |||
| a628f31859 | |||
| aae138a8cf | |||
| d2645589da | |||
| 6684a3426b | |||
| 7a8595ec87 | |||
| 06f76d918f | |||
| feb357ff17 | |||
| 34581e0d18 | |||
| 61c5b5e759 | |||
| c8eb780c67 | |||
| 4bdfa50b0c | |||
| 21008ccb43 | |||
| 362d327bf9 | |||
| 9a650fcc8f | |||
| 804a5a29ea | |||
| 6189fe23f1 | |||
| 7ba90e8696 | |||
| 0eed74fe10 | |||
| 466cc65069 | |||
| 9c41410f17 | |||
| bc71c524e0 | |||
| f46b2b4d86 | |||
| 9f17d63cdc | |||
| 92f8ec47db | |||
| a97484cd9b | |||
| ee6bc8c50e | |||
| c94ab84ab9 | |||
| 84a03215f4 | |||
| def93749eb | |||
| 945a1c0dfc | |||
| d752e096a3 | |||
| 880a3b81b0 | |||
| 191aface2e | |||
| e121dadb90 | |||
| bafcfbde3a | |||
| 98c463c62b | |||
| daf34d0f6c | |||
| bb505d461d | |||
| 00be0ac31e | |||
| fc5624a190 | |||
| 42ae3457b5 | |||
| 088027da7b | |||
| d37adee5df | |||
| c147e52d72 | |||
| d483258eef | |||
| 79c28103d5 | |||
| f36bf1437c | |||
| 4ad3aa2c06 | |||
| 357a4b165e | |||
| a15be253f5 | |||
| 572d1c16a6 | |||
| 39e22c07de | |||
| ca24573803 | |||
| 0cb3263a2e | |||
| fb2189c924 | |||
| 1383174b31 | |||
| c95bde7055 | |||
| 57f5e3a1a3 | |||
| 66ad52c199 | |||
| 2081dc9658 | |||
| b93255df3d | |||
| e1aef8700e | |||
| d5b3a39774 | |||
| 022e9e3807 | |||
| d333d463f6 | |||
| abc9f21ab9 | |||
| d216567c02 | |||
| 45bfd35a20 | |||
| 5a1b72387d | |||
| 66de566a00 | |||
| 18c2da7c7e | |||
| c4d807f1c4 | |||
| d1460ab9c7 | |||
| 9189c68800 | |||
| efbdabceca | |||
| 6e5a71f489 | |||
| e5c57f81ad | |||
| 0b1609c6cb | |||
| a10c68ef20 | |||
| 9aedeab406 | |||
| 137c34e3f5 | |||
| 25d29c305f | |||
| 9dcf9a1a43 | |||
| 2308b25bcf | |||
| b6c2159614 | |||
| 30c96a280d | |||
| 12c50df6a7 | |||
| c5124a01e6 | |||
| 6d57b49595 | |||
| d12c5bf2e1 | |||
| 42701e6c01 | |||
| d6c17aee79 | |||
| 17f8a06704 | |||
| e7b777890e | |||
| 5b03ce87ff | |||
| 2aebb9ed5e | |||
| e209fc689a |
@@ -0,0 +1,62 @@
|
||||
# 功能请求:在规则页面显示隧道倍率
|
||||
|
||||
## 问题描述
|
||||
|
||||
当前规则(Forward)页面在列表中显示隧道名称,但**不显示隧道的流量倍率(trafficRatio)**。管理员在管理规则时无法快速查看该规则所使用的隧道倍率信息,需要跳转到隧道页面才能查看。
|
||||
|
||||
## 期望行为
|
||||
|
||||
在规则列表页面中,在隧道名称旁边或单独列显示该隧道的流量倍率(例如:`1x`, `0.5x`, `2x`)。
|
||||
|
||||
## 建议实现位置
|
||||
|
||||
### 前端修改
|
||||
|
||||
1. **`vite-frontend/src/pages/forward.tsx`**
|
||||
- 在 `Forward` interface 中添加 `tunnelTrafficRatio?: number` 字段
|
||||
- 在表格列中添加倍率显示(可以在隧道名称 Chip 旁边或单独一列)
|
||||
- 从 `userTunnel` 或 `getTunnelList` API 获取隧道倍率信息
|
||||
|
||||
2. **显示格式建议**
|
||||
```tsx
|
||||
<Chip className="...">
|
||||
{forward.tunnelName} ({forward.tunnelTrafficRatio}x)
|
||||
</Chip>
|
||||
```
|
||||
或者单独一列:
|
||||
```tsx
|
||||
<TableCell>
|
||||
{forward.tunnelTrafficRatio}x
|
||||
</TableCell>
|
||||
```
|
||||
|
||||
### 后端修改
|
||||
|
||||
1. **`go-backend/internal/http/handler/handler.go`**
|
||||
- 在 `forwardList` 接口返回中添加隧道的 `trafficRatio` 字段
|
||||
- 需要在查询 Forward 时 JOIN Tunnel 表获取倍率信息
|
||||
|
||||
2. **或者在前端加载规则后,批量获取隧道信息**
|
||||
- 调用 `getTunnelList` 获取所有隧道信息
|
||||
- 根据 `tunnelId` 匹配倍率
|
||||
|
||||
## 相关文件
|
||||
|
||||
- 前端:`vite-frontend/src/pages/forward.tsx`
|
||||
- 前端类型:`vite-frontend/src/api/types.ts`
|
||||
- 后端:`go-backend/internal/http/handler/handler.go`
|
||||
- 隧道类型定义:`vite-frontend/src/api/types.ts` (TunnelApiItem)
|
||||
|
||||
## 优先级
|
||||
|
||||
中等 - 不影响核心功能,但能提升管理效率
|
||||
|
||||
## 截图参考
|
||||
|
||||
隧道页面已显示倍率:
|
||||
- 位置:隧道卡片统计信息区域
|
||||
- 显示格式:`流量倍率 {trafficRatio}x`
|
||||
|
||||
---
|
||||
|
||||
**Labels**: `enhancement`, `frontend`, `backend`, `ui/ux`
|
||||
@@ -0,0 +1,48 @@
|
||||
name: Publish Skill to npm
|
||||
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- 'v*'
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
publish:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
id-token: write
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Setup Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: '20'
|
||||
registry-url: 'https://registry.npmjs.org'
|
||||
|
||||
- name: Get version from tag
|
||||
id: version
|
||||
run: |
|
||||
if [ "${{ github.event_name }}" = "workflow_dispatch" ]; then
|
||||
VERSION=$(node -p "require('./skills/flvx-api/package.json').version")
|
||||
else
|
||||
VERSION="${GITHUB_REF#refs/tags/v}"
|
||||
fi
|
||||
echo "version=$VERSION" >> $GITHUB_OUTPUT
|
||||
echo "Publishing skill version: $VERSION"
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: skills/flvx-api
|
||||
run: npm publish --provenance --access public
|
||||
env:
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
|
||||
- name: Create GitHub Release
|
||||
if: github.event_name == 'push'
|
||||
uses: softprops/action-gh-release@v1
|
||||
with:
|
||||
name: Skill v${{ steps.version.outputs.version }}
|
||||
generate_release_notes: true
|
||||
files: skills/flvx-api/package.json
|
||||
@@ -62,6 +62,9 @@ go-gost/ss/
|
||||
.classpath
|
||||
.project
|
||||
.settings/
|
||||
|
||||
# OpenCode session metadata
|
||||
.entire/
|
||||
bin/
|
||||
tmp/
|
||||
*.swp
|
||||
@@ -259,6 +262,8 @@ gitee/
|
||||
doraemon.jks
|
||||
device.id
|
||||
commit.sh
|
||||
.opencode/
|
||||
analysis/
|
||||
sql/
|
||||
!go-backend/internal/store/sqlite/sql/
|
||||
!go-backend/internal/store/sqlite/sql/schema.sql
|
||||
@@ -266,3 +271,6 @@ sql/
|
||||
!go-backend/internal/store/postgres/sql/
|
||||
!go-backend/internal/store/postgres/sql/schema.sql
|
||||
!go-backend/internal/store/postgres/sql/data.sql
|
||||
go-backend/gost.db-shm
|
||||
.gitignore
|
||||
go-backend/gost.db-wal
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
# Plan: 搭建开发环境
|
||||
|
||||
## 目标
|
||||
为 Flux Panel 项目安装所有缺失的开发依赖,使 3 个子项目都能本地开发和构建。
|
||||
|
||||
## 当前状态
|
||||
|
||||
### ✅ 已安装
|
||||
| 工具 | 版本 | 用途 |
|
||||
|------|------|------|
|
||||
| Node.js | v20.19.2 | vite-frontend |
|
||||
| npm | 9.2.0 | vite-frontend |
|
||||
| Go | 1.24.4 | go-gost |
|
||||
| Docker | 29.1.4 | 容器化部署 |
|
||||
|
||||
### ❌ 缺失
|
||||
| 工具 | 需求版本 | 用途 |
|
||||
|------|----------|------|
|
||||
| Java | 21 | springboot-backend |
|
||||
| Maven | 3.x | 构建后端 |
|
||||
| Docker Compose | v2 | 容器编排 |
|
||||
|
||||
---
|
||||
|
||||
## 执行任务
|
||||
|
||||
### Task 1: 安装 Java 21
|
||||
```bash
|
||||
apt-get update && apt-get install -y openjdk-21-jdk
|
||||
```
|
||||
**验证**: `java -version` 应显示 openjdk 21
|
||||
|
||||
### Task 2: 安装 Maven
|
||||
```bash
|
||||
apt-get install -y maven
|
||||
```
|
||||
**验证**: `mvn -v` 应显示 Maven 3.x
|
||||
|
||||
### Task 3: 安装 Docker Compose Plugin
|
||||
```bash
|
||||
apt-get install -y docker-compose-plugin
|
||||
```
|
||||
**验证**: `docker compose version` 应显示版本号
|
||||
|
||||
### Task 4: 安装前端依赖
|
||||
```bash
|
||||
cd /root/flux-panel/vite-frontend && npm install
|
||||
```
|
||||
**验证**: `node_modules/` 目录存在
|
||||
|
||||
### Task 5: 验证后端可构建
|
||||
```bash
|
||||
cd /root/flux-panel/springboot-backend && mvn clean compile -q
|
||||
```
|
||||
**验证**: 编译成功无错误
|
||||
|
||||
### Task 6: 验证 Go 模块
|
||||
```bash
|
||||
cd /root/flux-panel/go-gost && go mod download
|
||||
```
|
||||
**验证**: 依赖下载成功
|
||||
|
||||
---
|
||||
|
||||
## 完成标准
|
||||
- [ ] `java -version` → openjdk 21
|
||||
- [ ] `mvn -v` → Maven 3.x
|
||||
- [ ] `docker compose version` → v2.x
|
||||
- [ ] 前端: `npm run dev` 可启动
|
||||
- [ ] 后端: `mvn compile` 成功
|
||||
- [ ] Go: `go build .` 成功
|
||||
@@ -1,9 +0,0 @@
|
||||
---
|
||||
active: true
|
||||
iteration: 1
|
||||
max_iterations: 100
|
||||
completion_promise: "DONE"
|
||||
started_at: "2026-02-15T16:38:58.212Z"
|
||||
session_id: "ses_39dd49703ffeveg711aA1D1YAk"
|
||||
---
|
||||
现状后端数据库兼容sqlite和postgresql,每一次新增功能需要维护两套数据库sql,需求是使用一个数据库驱动能同时兼容两个数据库,请仔细分析,列出计划,全量迁移,并且写好所有的测试,确保重构后所有的功能都能正常运行,由于工程量大,请写一个计划列表的markdown记录,每次完成一个就记录一下进度
|
||||
@@ -1,11 +1,12 @@
|
||||
# PROJECT KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Commit:** e5e22ba
|
||||
**Generated:** Tue Mar 24 2026
|
||||
**Commit:** 8ebde9d
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-rc10
|
||||
|
||||
## OVERVIEW
|
||||
FLVX (formerly Flux Panel) is a traffic forwarding management system built on a forked GOST v3 stack. It ships as a Go-based admin API (SQLite) + Vite/React UI + Go forwarding agent, with optional mobile WebView wrappers.
|
||||
FLVX (formerly Flux Panel) is a traffic forwarding management system built on a forked GOST v3 stack. It ships as a Go-based admin API (SQLite/PostgreSQL) + Vite/React UI + Go forwarding agent, with optional mobile WebView wrappers.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
@@ -13,12 +14,14 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
├── go-gost/ # Go forwarding agent (forked gost + local x/)
|
||||
│ └── x/ # Local fork of github.com/go-gost/x (replace => ./x)
|
||||
├── go-backend/ # Go Admin API (GORM + SQLite/PostgreSQL, net/http)
|
||||
├── vite-frontend/ # React/Vite dashboard (HeroUI + Tailwind)
|
||||
│ └── tests/contract/ # Integration/contract tests
|
||||
├── vite-frontend/ # React/Vite dashboard (shadcn bridge + Tailwind v4)
|
||||
│ └── src/shadcn-bridge/heroui/ # HeroUI-compatible facade
|
||||
├── docker-compose-v4.yml # Panel deploy (IPv4-only bridge)
|
||||
├── docker-compose-v6.yml # Panel deploy (IPv6-enabled bridge)
|
||||
├── panel_install.sh # Panel installer/upgrader (downloads compose)
|
||||
├── install.sh # Node installer/upgrader (downloads gost binary)
|
||||
└── .github/workflows/ # CI: build/push images + release artifacts
|
||||
└── .github/workflows/ # CI: build/test + Docker push + release artifacts
|
||||
```
|
||||
|
||||
## WHERE TO LOOK
|
||||
@@ -28,10 +31,15 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
| **Deploy (IPv6)** | `docker-compose-v6.yml` | Same as v4 + IPv6-enabled bridge |
|
||||
| **Panel install** | `panel_install.sh` | Picks v4/v6, generates `JWT_SECRET`, downloads compose |
|
||||
| **Node install** | `install.sh` | Installs `/etc/flux_agent/flux_agent` + writes `config.json`/`gost.json` + systemd `flux_agent.service` |
|
||||
| **Admin API** | `go-backend/` | Go Admin API (SQLite) |
|
||||
| **Web UI** | `vite-frontend/` | React/Vite dashboard (HeroUI + Tailwind) |
|
||||
| **Admin API** | `go-backend/` | Go Admin API (SQLite/PostgreSQL) |
|
||||
| **Web UI** | `vite-frontend/` | React/Vite dashboard (shadcn bridge + Tailwind v4) |
|
||||
| **UI Compatibility** | `vite-frontend/src/shadcn-bridge/heroui/` | HeroUI-compatible API wrappers backed by shadcn/radix |
|
||||
| **Theme Tokens** | `vite-frontend/src/styles/tailwind-theme.pcss` | Tailwind v4 `@theme inline` semantic color mapping |
|
||||
| **Go Agent** | `go-gost/` | Forwarding agent (forked gost + local x/) |
|
||||
| **Go Core** | `go-gost/x/` | Handlers/listeners/dialers + management API |
|
||||
| **Repository Layer** | `go-backend/internal/store/repo/` | GORM data access (repository.go 83k LOC) |
|
||||
| **Contract Tests** | `go-backend/tests/contract/` | Integration tests for auth, federation, tunnels |
|
||||
| **CI Workflows** | `.github/workflows/` | ci-build.yml, docker-build.yml, deploy-docs.yml |
|
||||
|
||||
## CODE MAP
|
||||
| Symbol | Type | Location | Role |
|
||||
@@ -40,13 +48,19 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
| `main` | Func | `go-backend/cmd/paneld/main.go` | Backend Entry |
|
||||
| `App` | Component | `vite-frontend/src/App.tsx` | Frontend Entry |
|
||||
| `main` | Func | `go-gost/main.go` | Agent Entry |
|
||||
|
||||
| `Repository` | Struct | `go-backend/internal/store/repo/repository.go` | Data Access Layer |
|
||||
| `Handler` | Struct | `go-backend/internal/http/handler/handler.go` | HTTP Handlers |
|
||||
| `websocket_reporter` | Func | `go-gost/x/socket/websocket_reporter.go` | Panel Telemetry |
|
||||
|
||||
## CONVENTIONS
|
||||
- **Skills & MCP**: Always prefer using available skills (via `skill` tool) and MCP tools when applicable. Check for relevant skills before implementing from scratch.
|
||||
- **Auth**: `Authorization` header carries the raw JWT token (no `Bearer` prefix) between `vite-frontend/` and `go-backend/`.
|
||||
- **Module Fork**: `go-gost/` uses `replace github.com/go-gost/x => ./x` and `go-gost/x/` is also its own Go module.
|
||||
- **Encryption**: Agent-to-panel communication uses AES encryption with node `secret` as PSK.
|
||||
- **API Envelope**: All REST responses follow `{code, msg, data, ts}` structure (code 0 = success).
|
||||
- **Frontend UI Layer**: Import UI primitives from `src/shadcn-bridge/heroui/*` (legacy-compatible facade), not direct `@heroui/*` packages.
|
||||
- **Tailwind v4 Semantic Colors**: `src/styles/globals.css` must import `src/styles/tailwind-theme.pcss`; removing it breaks semantic classes like `bg-primary`, `text-foreground`, and `border-input`.
|
||||
- **Go Versions**: `go-backend` uses Go 1.24, `go-gost` uses Go 1.23, `go-gost/x` uses Go 1.22.
|
||||
|
||||
## ANTI-PATTERNS (THIS PROJECT)
|
||||
- **DO NOT EDIT** generated protobuf output: `go-gost/x/internal/util/grpc/proto/*.pb.go`, `go-gost/x/internal/util/grpc/proto/*_grpc.pb.go`.
|
||||
@@ -54,6 +68,9 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
- **DO NOT MODIFY** `install.sh` or `panel_install.sh` locally - CI overwrites these on release.
|
||||
- **DO NOT** let backend handlers call `repo.DB()` directly — add a Repository method instead.
|
||||
- **DO NOT ADD** frontend tests - project has no test infrastructure (Vitest/Jest not configured).
|
||||
- **DO NOT REINTRODUCE** `@heroui/*` or `@nextui-org/*` dependencies; migration is now shadcn bridge-based.
|
||||
- **DO NOT** use `type:jsonb` or `type:serial` in GORM tags (SQLite incompatible).
|
||||
- **DO NOT** omit `TableName()` on new models — GORM pluralizes by default.
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
@@ -69,15 +86,21 @@ docker compose -f docker-compose-v6.yml up -d
|
||||
(cd go-backend && make build)
|
||||
(cd vite-frontend && npm run dev)
|
||||
(cd go-gost && go run .)
|
||||
|
||||
# Testing
|
||||
(cd go-backend && go test ./...)
|
||||
(cd go-backend && go test ./tests/contract/...)
|
||||
```
|
||||
|
||||
## UNIQUE STYLES
|
||||
- **Flat Monorepo**: Language-prefixed dirs (`go-backend`, `go-gost`, `vite-frontend`) instead of `apps/`/`libs/`.
|
||||
- **Asymmetric Go Layout**: `go-backend` follows `cmd/<app>/main.go` while `go-gost` uses `root/main.go`.
|
||||
- **Frontend Hybrid Mode**: `App.tsx` detects "H5 mode" (mobile WebView) vs desktop, dictating layout strategy.
|
||||
- **Experimental Bundler**: `vite-frontend` uses `rolldown-vite` (Rust-based) instead of standard Vite.
|
||||
- **Non-minified Builds**: `vite.config.ts` sets `minify: false`, `treeshake: false` for debugging.
|
||||
|
||||
## NOTES
|
||||
- LSP servers are not installed in this environment (gopls/jdtls/typescript-language-server); rely on grep-based navigation.
|
||||
- LSP servers are not installed in this environment (gopls/typescript-language-server); rely on grep-based navigation.
|
||||
- `vite-frontend/vite.config.ts` sets `minify: false` and disables treeshake; expect larger bundles.
|
||||
- `vite-frontend` uses `rolldown-vite` (experimental Rust bundler) instead of standard Vite.
|
||||
- Install scripts (`install.sh`, `panel_install.sh`) self-delete after execution - common pattern in one-liner installs.
|
||||
@@ -85,5 +108,16 @@ docker compose -f docker-compose-v6.yml up -d
|
||||
- CI dynamically injects `PINNED_VERSION` into install scripts and docker-compose files during releases.
|
||||
- `panel_install.sh` auto-detects IPv6 and modifies `/etc/docker/daemon.json` to enable IPv6 bridge.
|
||||
- Download proxy `https://gcode.hostcentral.cc/` used for GitHub downloads in China/restricted environments.
|
||||
- Backend has contract tests in `go-backend/tests/contract/` - frontend has no test infrastructure (Vitest/Jest not configured).
|
||||
- Backend has contract tests in `go-backend/tests/contract/` - frontend has no test infrastructure.
|
||||
- `analysis/3x-ui/` contains a separate git repo for reference/comparison - not part of FLVX core.
|
||||
- CI workflows: `ci-build.yml` (build check), `docker-build.yml` (multi-arch images + release), `deploy-docs.yml` (MkDocs).
|
||||
- PostgreSQL migration supported via `panel_install.sh` menu option using pgloader.
|
||||
- Repository layer is large: `repository.go` (83k LOC), `repository_mutations.go` (43k LOC).
|
||||
- Button visual parity relies on `vite-frontend/src/shadcn-bridge/heroui/button.tsx` color mapping + `vite-frontend/src/styles/tailwind-theme.pcss` token export.
|
||||
|
||||
## PLAN DOCUMENT RULE
|
||||
- Every new implementation plan must have a dedicated Markdown plan document.
|
||||
- Store plan documents under `plans/`.
|
||||
- Use an incrementing numeric prefix and a short plan-summary name: `NNN-<plan-summary>.md` (for example, `001-auth-refactor.md`, `002-federation-api-cleanup.md`).
|
||||
- The numeric prefix must increase by 1 for each new plan.
|
||||
- In each plan document, keep a task checklist and mark each task as completed immediately after finishing it.
|
||||
|
||||
@@ -33,13 +33,13 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -
|
||||
#### 安装特定版本
|
||||
从 [Releases](https://github.com/Sagit-chu/flux-panel/releases) 页面复制对应版本的安装命令,脚本会自动安装该版本而非最新版。
|
||||
|
||||
面板端(以 2.1.0 为例):
|
||||
面板端(以 2.1.9-beta6 为例):
|
||||
```bash
|
||||
curl -L https://github.com/Sagit-chu/flux-panel/releases/download/2.1.0/panel_install.sh -o panel_install.sh && chmod +x panel_install.sh && ./panel_install.sh
|
||||
curl -L https://github.com/Sagit-chu/flux-panel/releases/download/2.1.9-beta6/panel_install.sh -o panel_install.sh && chmod +x panel_install.sh && ./panel_install.sh
|
||||
```
|
||||
节点端(以 2.1.0 为例):
|
||||
节点端(以 2.1.9-beta6 为例):
|
||||
```bash
|
||||
curl -L https://github.com/Sagit-chu/flux-panel/releases/download/2.1.0/install.sh -o install.sh && chmod +x install.sh && ./install.sh
|
||||
curl -L https://github.com/Sagit-chu/flux-panel/releases/download/2.1.9-beta6/install.sh -o install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
#### PostgreSQL 部署(Docker Compose)
|
||||
@@ -129,27 +129,28 @@ docker compose up -d
|
||||
- **License**: Apache License 2.0
|
||||
|
||||
## Modifications
|
||||
The following major changes and additions have been made in this fork (FLVX):
|
||||
This fork (FLVX) is no longer a light patch on top of the upstream project. It has been deeply reworked, with both backend and frontend rebuilt around a Go-based architecture.
|
||||
|
||||
### 1. Backend Architecture (Replaced)
|
||||
- **Removed**: The original `springboot-backend/` (Java/Spring Boot) has been entirely removed.
|
||||
- **Added**: A new `go-backend/` (Go/SQLite) implementation replaces the original backend.
|
||||
### 1. Backend (Rewritten)
|
||||
- **Removed**: The original `springboot-backend/` (Java/Spring Boot) implementation.
|
||||
- **Added**: A fully rewritten `go-backend/` service (Go), including updated data and API handling for panel management.
|
||||
|
||||
### 2. Forwarding Agent (Modified)
|
||||
- **Modified**: `go-gost/` - Modified forwarding agent wrapper.
|
||||
- **Modified**: `go-gost/x/` - Modified local fork of the `gost` extensions library.
|
||||
### 2. Frontend (Reworked)
|
||||
- **Reworked**: `vite-frontend/` has been substantially rebuilt to match the new backend contract and current UI layer architecture.
|
||||
- **Updated**: Dashboard pages/components and interaction flows for the current React/Vite stack.
|
||||
|
||||
### 3. Frontend (Modified)
|
||||
- **Modified**: `vite-frontend/` - Significant updates to the React/Vite dashboard to compatible with the new Go backend, including UI/UX improvements (HeroUI + Tailwind).
|
||||
### 3. Forwarding Stack (Modified)
|
||||
- **Modified**: `go-gost/` forwarding agent wrapper.
|
||||
- **Modified**: `go-gost/x/` local fork of `github.com/go-gost/x`.
|
||||
|
||||
### 4. Mobile Applications (Removed)
|
||||
- **Removed**: `android-app/` - Source code for the Android client.
|
||||
- **Removed**: `ios-app/` - Source code for the iOS client.
|
||||
### 4. Mobile Clients (Removed)
|
||||
- **Removed**: `android-app/` source code.
|
||||
- **Removed**: `ios-app/` source code.
|
||||
|
||||
### 5. Infrastructure & Scripts
|
||||
- **Modified**: `docker-compose.yml` (installer output name, auto-selects IPv4/IPv6 template, updated for Go backend).
|
||||
- **Modified**: `install.sh`, `panel_install.sh` (Updated installation logic).
|
||||
- **Added**: `AGENTS.md` (Project documentation).
|
||||
### 5. Deployment & Project Infrastructure
|
||||
- **Updated**: Docker deployment templates and installer output flow (IPv4/IPv6 compose variants).
|
||||
- **Updated**: Release installation scripts (`install.sh`, `panel_install.sh`) and supporting automation.
|
||||
- **Added/Updated**: Project-level engineering documentation (for example `AGENTS.md`).
|
||||
|
||||
---
|
||||
|
||||
|
||||
+220
@@ -0,0 +1,220 @@
|
||||
# AI Skill 使用指南
|
||||
|
||||
让大模型直接操作 FLVX 面板的技能包。支持 OpenCode、OpenClaw、Claude Code 等工具。
|
||||
|
||||
## 安装
|
||||
|
||||
### 方式 1: npm (推荐)
|
||||
|
||||
```bash
|
||||
npm install -g @flvx/skill-api
|
||||
```
|
||||
|
||||
postinstall 脚本会自动链接到 `~/.agents/skills/flvx-api/`。
|
||||
|
||||
### 方式 2: 手动链接
|
||||
|
||||
```bash
|
||||
# 从 FLVX 源码
|
||||
cd /path/to/flvx
|
||||
mkdir -p ~/.agents/skills
|
||||
ln -sf $(pwd)/skills/flvx-api ~/.agents/skills/
|
||||
|
||||
# 或从 GitHub
|
||||
git clone https://github.com/Sagit-chu/flvx.git
|
||||
cd flvx
|
||||
ln -sf $(pwd)/skills/flvx-api ~/.agents/skills/
|
||||
```
|
||||
|
||||
## 配置
|
||||
|
||||
设置环境变量:
|
||||
|
||||
```bash
|
||||
export FLVX_BASE_URL="https://your-panel.example.com"
|
||||
export FLVX_USERNAME="admin"
|
||||
export FLVX_PASSWORD="your-password"
|
||||
```
|
||||
|
||||
或使用凭证文件:
|
||||
|
||||
```bash
|
||||
mkdir -p ~/.flvx
|
||||
cat > ~/.flvx/.env << 'EOF'
|
||||
export FLVX_BASE_URL="https://panel.example.com"
|
||||
export FLVX_USERNAME="admin"
|
||||
export FLVX_PASSWORD="your-password"
|
||||
EOF
|
||||
chmod 600 ~/.flvx/.env
|
||||
source ~/.flvx/.env
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 工具接入方法
|
||||
|
||||
### OpenCode
|
||||
|
||||
OpenCode 是命令行 AI 编程助手,支持通过 skills 扩展能力。
|
||||
|
||||
**安装 skill:**
|
||||
```bash
|
||||
npm install -g @flvx/skill-api
|
||||
```
|
||||
|
||||
**使用:**
|
||||
```bash
|
||||
export FLVX_BASE_URL="https://panel.example.com"
|
||||
export FLVX_USERNAME="admin"
|
||||
export FLVX_PASSWORD="your-password"
|
||||
|
||||
opencode
|
||||
```
|
||||
|
||||
**示例对话:**
|
||||
```
|
||||
你: 查看我的转发列表
|
||||
你: 创建一个转发到 192.168.1.100:80 使用隧道 1
|
||||
你: 检查节点状态
|
||||
你: 查看流量使用情况
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### OpenClaw
|
||||
|
||||
OpenClaw 同样支持 skills 机制。
|
||||
|
||||
**安装 skill:**
|
||||
```bash
|
||||
npm install -g @flvx/skill-api
|
||||
|
||||
# 或手动链接
|
||||
mkdir -p ~/.openclaw/skills
|
||||
ln -sf /path/to/flvx/skills/flvx-api ~/.openclaw/skills/flvx-api
|
||||
```
|
||||
|
||||
**使用:**
|
||||
```bash
|
||||
openclaw
|
||||
|
||||
>>> 查看所有节点状态
|
||||
>>> 给用户 alice 分配 50GB 流量
|
||||
>>> 导出系统备份
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Claude Code
|
||||
|
||||
Claude Code 是 Anthropic 官方的命令行工具,支持通过 CLAUDE.md 扩展。
|
||||
|
||||
#### 方式 1: 项目级 CLAUDE.md
|
||||
|
||||
在项目根目录创建 `CLAUDE.md`:
|
||||
|
||||
```markdown
|
||||
# FLVX API Skill
|
||||
|
||||
你可以通过 REST API 操作 FLVX 面板。
|
||||
|
||||
## 环境变量
|
||||
- FLVX_BASE_URL: 面板地址
|
||||
- FLVX_USERNAME: 用户名
|
||||
- FLVX_PASSWORD: 密码
|
||||
|
||||
## 认证规则
|
||||
- Authorization 头使用原始 JWT token,不加 "Bearer " 前缀
|
||||
- 所有 API 使用 POST 方法
|
||||
|
||||
## 常用 API
|
||||
|
||||
### 登录获取 token
|
||||
POST /api/v1/user/login
|
||||
{"username": "...", "password": "..."}
|
||||
|
||||
### 查看转发列表
|
||||
POST /api/v1/forward/list
|
||||
Authorization: <token>
|
||||
{}
|
||||
|
||||
### 创建转发
|
||||
POST /api/v1/forward/create
|
||||
{"name": "xxx", "tunnelId": 1, "remoteAddr": "1.2.3.4:80"}
|
||||
|
||||
### 查看节点
|
||||
POST /api/v1/node/list
|
||||
{}
|
||||
```
|
||||
|
||||
**使用:**
|
||||
```bash
|
||||
cd /path/to/your/project
|
||||
claude
|
||||
```
|
||||
|
||||
#### 方式 2: 全局 CLAUDE.md
|
||||
|
||||
```bash
|
||||
mkdir -p ~/.claude
|
||||
cat > ~/.claude/CLAUDE.md << 'EOF'
|
||||
# FLVX Panel Operations
|
||||
|
||||
使用 FLVX REST API 操作流量转发面板。
|
||||
|
||||
环境变量: FLVX_BASE_URL, FLVX_USERNAME, FLVX_PASSWORD
|
||||
调用方式: curl -X POST "$FLVX_BASE_URL/api/v1/..." -H "Authorization: $TOKEN"
|
||||
注意: Authorization 不要加 Bearer 前缀
|
||||
EOF
|
||||
```
|
||||
|
||||
#### 方式 3: 复制 SKILL.md
|
||||
|
||||
```bash
|
||||
cat ~/.agents/skills/flvx-api/SKILL.md >> ~/.claude/CLAUDE.md
|
||||
```
|
||||
|
||||
**示例对话:**
|
||||
```
|
||||
>>> 帮我查看 FLVX 面板上有哪些节点
|
||||
>>> 创建一个名为 test 的转发,目标地址 10.0.0.1:80
|
||||
>>> 查看我的流量使用情况
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## API 覆盖
|
||||
|
||||
| 模块 | 操作 |
|
||||
|------|------|
|
||||
| 认证 | 登录、Token 管理 |
|
||||
| 用户 | 增删改查、流量重置、密码 |
|
||||
| 节点 | 增删改查、安装、升级、状态 |
|
||||
| 隧道 | 增删改查、用户分配 |
|
||||
| 转发 | 增删改查、暂停/恢复、诊断 |
|
||||
| 分组 | 用户/隧道分组、权限 |
|
||||
| 限速 | 增删改查 |
|
||||
| 联邦 | 节点共享、远程节点 |
|
||||
| 备份 | 导出/导入 |
|
||||
|
||||
## 安全提示
|
||||
|
||||
- ⚠️ 环境变量在进程列表中可见
|
||||
- 使用 `~/.flvx/.env` 文件并设置 `chmod 600`
|
||||
- 添加 `export HISTIGNORE="*FLVX_PASSWORD*"` 防止密码进入历史记录
|
||||
- Token 仅在会话内存中缓存,不写入磁盘
|
||||
|
||||
## 发布
|
||||
|
||||
维护者可通过以下方式发布新版本:
|
||||
|
||||
```bash
|
||||
# 方式 1: 推送 tag
|
||||
git tag skill-v2.1.6
|
||||
git push --tags
|
||||
|
||||
# 方式 2: GitHub Actions 手动触发
|
||||
# 在 Actions 页面运行 publish-skill workflow
|
||||
```
|
||||
|
||||
需要在 GitHub 仓库设置 `NPM_TOKEN` secret。
|
||||
@@ -18,6 +18,7 @@
|
||||
- [安装部署](./install.md)
|
||||
- [使用指南](./usage.md)
|
||||
- [PostgreSQL 数据库指南](./postgresql.md)
|
||||
- [AI Skill 接入](./ai-skill.md) - 让大模型直接操作面板
|
||||
- [常见问题](./faq.md)
|
||||
|
||||
## 免责声明
|
||||
|
||||
+17
-8
@@ -1,8 +1,13 @@
|
||||
# GO BACKEND KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Go-based Admin API for FLVX. Replaced legacy Spring Boot backend.
|
||||
**Stack:** Go 1.23, net/http (std lib), GORM + SQLite/PostgreSQL (glebarez/sqlite - CGO-free).
|
||||
**Stack:** Go 1.24, net/http (std lib), GORM + SQLite/PostgreSQL (glebarez/sqlite - CGO-free).
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
@@ -17,14 +22,15 @@ go-backend/
|
||||
│ ├── store/
|
||||
│ │ ├── model/model.go # GORM model structs (single source of truth)
|
||||
│ │ └── repo/ # Data Access Layer (Repository pattern, GORM)
|
||||
│ │ ├── repository.go # Core queries, Open/OpenPostgres, AutoMigrate
|
||||
│ │ ├── repository_mutations.go # Mutation helpers (user/node/tunnel/forward CRUD)
|
||||
│ │ ├── repository_federation.go# Federation-specific queries
|
||||
│ │ ├── repository.go # Core queries, Open/OpenPostgres, AutoMigrate (83k LOC)
|
||||
│ │ ├── repository_mutations.go # Mutation helpers (user/node/tunnel/forward CRUD, 43k LOC)
|
||||
│ │ ├── repository_federation.go # Federation-specific queries
|
||||
│ │ ├── repository_flow.go # Flow/forward status queries
|
||||
│ │ └── repository_control.go # Control plane queries
|
||||
│ │ ├── repository_control.go # Control plane queries
|
||||
│ │ └── repository_groups.go # Group management queries
|
||||
│ └── auth/ # Auth logic
|
||||
├── tests/ # Integration/Contract tests
|
||||
├── Dockerfile # Multi-stage build (alpine)
|
||||
├── tests/contract/ # Integration/contract tests (14 tests)
|
||||
├── Dockerfile # Multi-stage build (golang:1.24-bookworm → debian:bookworm-slim)
|
||||
└── Makefile # Build commands
|
||||
```
|
||||
|
||||
@@ -36,6 +42,7 @@ go-backend/
|
||||
| **Repository** | `go-backend/internal/store/repo/` | GORM-based queries, all DB ops encapsulated |
|
||||
| **Auth Middleware** | `go-backend/internal/http/middleware/jwt.go` | Extracts `Authorization` header |
|
||||
| **WebSocket** | `go-backend/internal/ws/` | Real-time updates (traffic, status) |
|
||||
| **Contract Tests** | `go-backend/tests/contract/` | Integration tests for auth, federation, tunnels |
|
||||
|
||||
## CONVENTIONS
|
||||
- **GORM ORM**: Uses GORM with `glebarez/sqlite` (CGO-free) and `gorm.io/driver/postgres`.
|
||||
@@ -47,6 +54,7 @@ go-backend/
|
||||
- **API Envelope**: All responses use `response.R{code, msg, data, ts}` structure.
|
||||
- **Config**: Loaded from environment variables (see `cmd/paneld/main.go`).
|
||||
- **SQLite Constraints**: `MaxOpenConns(1)`, WAL mode, busy_timeout=5000.
|
||||
- **PostgreSQL**: Supported via `DB_TYPE=postgres` and `DATABASE_URL` env vars.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- **DO NOT** let handlers call `repo.DB()` directly — add a Repository method instead.
|
||||
@@ -58,6 +66,7 @@ go-backend/
|
||||
```bash
|
||||
cd go-backend
|
||||
go run ./cmd/paneld # Default: SERVER_ADDR=:6365
|
||||
go test ./...
|
||||
go test ./... # Unit tests
|
||||
go test ./tests/contract/... # Contract tests
|
||||
make build
|
||||
```
|
||||
|
||||
@@ -49,7 +49,7 @@ func New(cfg config.Config) (*App, error) {
|
||||
Handler: router,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
ReadHeaderTimeout: 5 * time.Second,
|
||||
WriteTimeout: 30 * time.Second,
|
||||
WriteTimeout: 2 * time.Minute,
|
||||
IdleTimeout: 60 * time.Second,
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,402 @@
|
||||
package health
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
)
|
||||
|
||||
type nodeCommander interface {
|
||||
SendCommand(nodeID int64, cmdType string, data interface{}, timeout time.Duration) (ws.CommandResult, error)
|
||||
}
|
||||
|
||||
const serviceMonitorReportInterval = 30 * time.Second // DB write interval per monitor
|
||||
|
||||
type Checker struct {
|
||||
repo *repo.Repository
|
||||
commander nodeCommander
|
||||
lastRun map[int64]int64
|
||||
inFlight map[int64]struct{}
|
||||
|
||||
// In-memory latest result per monitor (for real-time API reads)
|
||||
latestResults map[int64]*model.ServiceMonitorResult
|
||||
lastDBWrite map[int64]int64 // last DB write timestamp per monitorID
|
||||
|
||||
mu sync.RWMutex
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
checking int32 // atomic flag: 1 = runChecks running, 0 = idle
|
||||
}
|
||||
|
||||
func NewChecker(repo *repo.Repository, commander nodeCommander) *Checker {
|
||||
return &Checker{
|
||||
repo: repo,
|
||||
commander: commander,
|
||||
lastRun: make(map[int64]int64),
|
||||
inFlight: make(map[int64]struct{}),
|
||||
latestResults: make(map[int64]*model.ServiceMonitorResult),
|
||||
lastDBWrite: make(map[int64]int64),
|
||||
}
|
||||
}
|
||||
|
||||
// GetLatestCached returns the in-memory latest results (updated every 1s).
|
||||
// Returns nil if no results are cached.
|
||||
func (c *Checker) GetLatestCached() []*model.ServiceMonitorResult {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
results := make([]*model.ServiceMonitorResult, 0, len(c.latestResults))
|
||||
for _, r := range c.latestResults {
|
||||
results = append(results, r)
|
||||
}
|
||||
return results
|
||||
}
|
||||
|
||||
func (c *Checker) Start(ctx context.Context) {
|
||||
c.mu.Lock()
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
c.cancel = cancel
|
||||
c.mu.Unlock()
|
||||
|
||||
c.runChecks(ctx)
|
||||
|
||||
for {
|
||||
limits := c.loadServiceMonitorLimits()
|
||||
scanInterval := time.Duration(limits.CheckerScanIntervalSec) * time.Second
|
||||
if scanInterval <= 0 {
|
||||
scanInterval = 1 * time.Second
|
||||
}
|
||||
|
||||
timer := time.NewTimer(scanInterval)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
timer.Stop()
|
||||
return
|
||||
case <-timer.C:
|
||||
c.runChecks(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Checker) Stop() {
|
||||
c.mu.Lock()
|
||||
if c.cancel != nil {
|
||||
c.cancel()
|
||||
}
|
||||
c.mu.Unlock()
|
||||
c.wg.Wait()
|
||||
}
|
||||
|
||||
func (c *Checker) RunOnce(m *model.ServiceMonitor) (*model.ServiceMonitorResult, error) {
|
||||
if c == nil {
|
||||
return nil, errors.New("checker not initialized")
|
||||
}
|
||||
if m == nil {
|
||||
return nil, errors.New("monitor is nil")
|
||||
}
|
||||
limits := c.loadServiceMonitorLimits()
|
||||
return c.executeCheck(m, time.Now().UnixMilli(), limits), nil
|
||||
}
|
||||
|
||||
func (c *Checker) runChecks(ctx context.Context) {
|
||||
// Skip if previous round is still running (interval < timeout guard)
|
||||
if !atomic.CompareAndSwapInt32(&c.checking, 0, 1) {
|
||||
return
|
||||
}
|
||||
defer atomic.StoreInt32(&c.checking, 0)
|
||||
|
||||
if c == nil || c.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
limits := c.loadServiceMonitorLimits()
|
||||
monitors, err := c.repo.ListEnabledServiceMonitors()
|
||||
if err != nil {
|
||||
log.Printf("service monitor scheduler failed op=list_enabled err=%v", err)
|
||||
return
|
||||
}
|
||||
if len(monitors) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Use persisted result timestamps to avoid restart bursts.
|
||||
latest, err := c.repo.GetLatestServiceMonitorResults()
|
||||
if err != nil {
|
||||
log.Printf("service monitor scheduler failed op=get_latest_results err=%v", err)
|
||||
latest = nil
|
||||
}
|
||||
persistedLast := make(map[int64]int64, len(latest))
|
||||
for _, r := range latest {
|
||||
if r.MonitorID <= 0 || r.Timestamp <= 0 {
|
||||
continue
|
||||
}
|
||||
persistedLast[r.MonitorID] = r.Timestamp
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
due := make([]model.ServiceMonitor, 0, len(monitors))
|
||||
for _, m := range monitors {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
intervalSec := m.IntervalSec
|
||||
if intervalSec <= 0 {
|
||||
intervalSec = limits.DefaultIntervalSec
|
||||
}
|
||||
if intervalSec < limits.MinIntervalSec {
|
||||
intervalSec = limits.MinIntervalSec
|
||||
}
|
||||
intervalMs := int64(intervalSec) * 1000
|
||||
|
||||
c.mu.Lock()
|
||||
if _, ok := c.inFlight[m.ID]; ok {
|
||||
c.mu.Unlock()
|
||||
continue
|
||||
}
|
||||
|
||||
lastSeen := persistedLast[m.ID]
|
||||
if v := c.lastRun[m.ID]; v > lastSeen {
|
||||
lastSeen = v
|
||||
}
|
||||
if lastSeen > 0 && intervalMs > 0 && now-lastSeen < intervalMs {
|
||||
c.mu.Unlock()
|
||||
continue
|
||||
}
|
||||
|
||||
c.inFlight[m.ID] = struct{}{}
|
||||
// Use now as a best-effort guard against overlapping scans; the final
|
||||
// timestamp is updated again when the result is persisted.
|
||||
c.lastRun[m.ID] = now
|
||||
c.mu.Unlock()
|
||||
|
||||
due = append(due, m)
|
||||
}
|
||||
if len(due) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
workerLimit := limits.WorkerLimit
|
||||
if workerLimit <= 0 {
|
||||
workerLimit = 1
|
||||
}
|
||||
if workerLimit > len(due) {
|
||||
workerLimit = len(due)
|
||||
}
|
||||
|
||||
jobs := make(chan model.ServiceMonitor, len(due))
|
||||
for _, m := range due {
|
||||
jobs <- m
|
||||
}
|
||||
close(jobs)
|
||||
|
||||
reportIntervalMs := int64(serviceMonitorReportInterval / time.Millisecond)
|
||||
|
||||
for i := 0; i < workerLimit; i++ {
|
||||
c.wg.Add(1)
|
||||
go func() {
|
||||
defer c.wg.Done()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case m, ok := <-jobs:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
ts := time.Now().UnixMilli()
|
||||
result := c.executeCheck(&m, ts, limits)
|
||||
|
||||
// Always update in-memory cache for real-time reads
|
||||
c.mu.Lock()
|
||||
c.latestResults[m.ID] = result
|
||||
c.lastRun[m.ID] = result.Timestamp
|
||||
delete(c.inFlight, m.ID)
|
||||
|
||||
// Only write to DB every 30s per monitor
|
||||
lastWrite := c.lastDBWrite[m.ID]
|
||||
writeToDB := ts-lastWrite >= reportIntervalMs
|
||||
if writeToDB {
|
||||
c.lastDBWrite[m.ID] = ts
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
if writeToDB {
|
||||
if err := c.repo.InsertServiceMonitorResult(result); err != nil {
|
||||
log.Printf("monitoring write failed op=service_monitor_result.insert monitor_id=%d err=%v", result.MonitorID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Checker) executeCheck(m *model.ServiceMonitor, timestamp int64, limits monitoring.ServiceMonitorLimits) *model.ServiceMonitorResult {
|
||||
result := &model.ServiceMonitorResult{
|
||||
MonitorID: m.ID,
|
||||
NodeID: m.NodeID,
|
||||
Timestamp: timestamp,
|
||||
}
|
||||
|
||||
timeoutSec := m.TimeoutSec
|
||||
if timeoutSec <= 0 {
|
||||
timeoutSec = limits.DefaultTimeoutSec
|
||||
}
|
||||
if timeoutSec < limits.MinTimeoutSec {
|
||||
timeoutSec = limits.MinTimeoutSec
|
||||
}
|
||||
if timeoutSec > limits.MaxTimeoutSec {
|
||||
timeoutSec = limits.MaxTimeoutSec
|
||||
}
|
||||
|
||||
timeout := time.Duration(timeoutSec) * time.Second
|
||||
|
||||
// When nodeId is set, run checks on the specified node.
|
||||
if m.NodeID > 0 {
|
||||
c.checkOnNode(m, timeoutSec, timeout, result)
|
||||
return result
|
||||
}
|
||||
|
||||
switch strings.ToLower(strings.TrimSpace(m.Type)) {
|
||||
case "tcp":
|
||||
c.checkTCP(m.Target, timeout, result)
|
||||
case "icmp":
|
||||
result.Success = 0
|
||||
result.ErrorMessage = "ICMP 监控必须指定执行节点"
|
||||
default:
|
||||
result.Success = 0
|
||||
result.ErrorMessage = fmt.Sprintf("不支持的检查类型: %s", m.Type)
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func (c *Checker) loadServiceMonitorLimits() monitoring.ServiceMonitorLimits {
|
||||
defaults := monitoring.DefaultServiceMonitorLimits()
|
||||
if c == nil || c.repo == nil {
|
||||
return defaults
|
||||
}
|
||||
cfg, err := c.repo.GetConfigsByNames([]string{
|
||||
monitoring.ConfigServiceMonitorCheckerScanIntervalSec,
|
||||
monitoring.ConfigServiceMonitorWorkerLimit,
|
||||
monitoring.ConfigServiceMonitorMinIntervalSec,
|
||||
monitoring.ConfigServiceMonitorDefaultIntervalSec,
|
||||
monitoring.ConfigServiceMonitorMinTimeoutSec,
|
||||
monitoring.ConfigServiceMonitorDefaultTimeoutSec,
|
||||
monitoring.ConfigServiceMonitorMaxTimeoutSec,
|
||||
})
|
||||
if err != nil {
|
||||
return defaults
|
||||
}
|
||||
return monitoring.ServiceMonitorLimitsFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
type serviceMonitorCheckRequest struct {
|
||||
MonitorID int64 `json:"monitorId"`
|
||||
Type string `json:"type"`
|
||||
Target string `json:"target"`
|
||||
TimeoutSec int `json:"timeoutSec"`
|
||||
}
|
||||
|
||||
func (c *Checker) checkOnNode(m *model.ServiceMonitor, timeoutSec int, timeout time.Duration, result *model.ServiceMonitorResult) {
|
||||
if c == nil || m == nil || result == nil {
|
||||
return
|
||||
}
|
||||
if c.commander == nil {
|
||||
result.Success = 0
|
||||
result.ErrorMessage = "节点检查不可用"
|
||||
return
|
||||
}
|
||||
|
||||
checkType := strings.ToLower(strings.TrimSpace(m.Type))
|
||||
if checkType != "tcp" && checkType != "icmp" {
|
||||
result.Success = 0
|
||||
result.ErrorMessage = fmt.Sprintf("不支持的检查类型: %s", m.Type)
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(m.Target) == "" {
|
||||
result.Success = 0
|
||||
result.ErrorMessage = "检查目标为空"
|
||||
return
|
||||
}
|
||||
|
||||
req := serviceMonitorCheckRequest{
|
||||
MonitorID: m.ID,
|
||||
Type: checkType,
|
||||
Target: m.Target,
|
||||
TimeoutSec: timeoutSec,
|
||||
}
|
||||
|
||||
cmdTimeout := timeout
|
||||
if cmdTimeout < 2*time.Second {
|
||||
cmdTimeout = 2 * time.Second
|
||||
}
|
||||
cmdTimeout = cmdTimeout + 2*time.Second
|
||||
|
||||
cmdRes, err := c.commander.SendCommand(m.NodeID, "ServiceMonitorCheck", req, cmdTimeout)
|
||||
if err != nil {
|
||||
result.Success = 0
|
||||
result.ErrorMessage = err.Error()
|
||||
return
|
||||
}
|
||||
if cmdRes.Data == nil {
|
||||
result.Success = 0
|
||||
result.ErrorMessage = "节点返回为空"
|
||||
return
|
||||
}
|
||||
|
||||
if v, ok := cmdRes.Data["success"]; ok {
|
||||
if b, ok := v.(bool); ok {
|
||||
if b {
|
||||
result.Success = 1
|
||||
} else {
|
||||
result.Success = 0
|
||||
}
|
||||
}
|
||||
}
|
||||
if v, ok := cmdRes.Data["latencyMs"]; ok {
|
||||
if f, ok := v.(float64); ok {
|
||||
result.LatencyMs = f
|
||||
}
|
||||
}
|
||||
if v, ok := cmdRes.Data["statusCode"]; ok {
|
||||
if f, ok := v.(float64); ok {
|
||||
result.StatusCode = int(f)
|
||||
}
|
||||
}
|
||||
if v, ok := cmdRes.Data["errorMessage"]; ok {
|
||||
if s, ok := v.(string); ok {
|
||||
result.ErrorMessage = s
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Checker) checkTCP(target string, timeout time.Duration, result *model.ServiceMonitorResult) {
|
||||
start := time.Now()
|
||||
|
||||
conn, err := net.DialTimeout("tcp", target, timeout)
|
||||
latency := time.Since(start)
|
||||
|
||||
result.LatencyMs = float64(latency.Milliseconds())
|
||||
|
||||
if err != nil {
|
||||
result.Success = 0
|
||||
result.ErrorMessage = err.Error()
|
||||
return
|
||||
}
|
||||
_ = conn.Close()
|
||||
result.Success = 1
|
||||
}
|
||||
@@ -0,0 +1,477 @@
|
||||
package health
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
)
|
||||
|
||||
type fakeCommander struct {
|
||||
lastNodeID int64
|
||||
lastType string
|
||||
lastData interface{}
|
||||
res ws.CommandResult
|
||||
err error
|
||||
}
|
||||
|
||||
type delayedCommander struct {
|
||||
delayByMonitorID map[int64]time.Duration
|
||||
}
|
||||
|
||||
func (d *delayedCommander) SendCommand(nodeID int64, cmdType string, data interface{}, _ time.Duration) (ws.CommandResult, error) {
|
||||
_ = nodeID
|
||||
_ = cmdType
|
||||
if req, ok := data.(serviceMonitorCheckRequest); ok {
|
||||
if delay := d.delayByMonitorID[req.MonitorID]; delay > 0 {
|
||||
time.Sleep(delay)
|
||||
}
|
||||
}
|
||||
return ws.CommandResult{
|
||||
Success: true,
|
||||
Data: map[string]interface{}{
|
||||
"success": true,
|
||||
"latencyMs": float64(1),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeCommander) SendCommand(nodeID int64, cmdType string, data interface{}, _ time.Duration) (ws.CommandResult, error) {
|
||||
f.lastNodeID = nodeID
|
||||
f.lastType = cmdType
|
||||
f.lastData = data
|
||||
return f.res, f.err
|
||||
}
|
||||
|
||||
func TestTCPHealthCheckViaMonitor(t *testing.T) {
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
addr := listener.Addr().String()
|
||||
|
||||
go func() {
|
||||
for {
|
||||
conn, err := listener.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
conn.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
t.Run("successful tcp check", func(t *testing.T) {
|
||||
checker := NewChecker(nil, nil)
|
||||
limits := checker.loadServiceMonitorLimits()
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
Type: "tcp",
|
||||
Target: addr,
|
||||
TimeoutSec: 5,
|
||||
}
|
||||
result := checker.executeCheck(monitor, now, limits)
|
||||
if result.Success != 1 {
|
||||
t.Fatalf("expected success, got error: %s", result.ErrorMessage)
|
||||
}
|
||||
if result.LatencyMs < 0 {
|
||||
t.Fatalf("expected non-negative latency, got %f", result.LatencyMs)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("failed tcp check - connection refused", func(t *testing.T) {
|
||||
checker := NewChecker(nil, nil)
|
||||
limits := checker.loadServiceMonitorLimits()
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
Type: "tcp",
|
||||
Target: "127.0.0.1:1",
|
||||
TimeoutSec: 1,
|
||||
}
|
||||
result := checker.executeCheck(monitor, now, limits)
|
||||
if result.Success == 1 {
|
||||
t.Fatalf("expected failure for connection refused")
|
||||
}
|
||||
if result.ErrorMessage == "" {
|
||||
t.Fatalf("expected error message")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestCheckerRunChecks(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
tcpAddr := listener.Addr().String()
|
||||
|
||||
go func() {
|
||||
for {
|
||||
conn, err := listener.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
conn.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
monitors := []*model.ServiceMonitor{
|
||||
{
|
||||
Name: "TCP Monitor",
|
||||
Type: "tcp",
|
||||
Target: tcpAddr,
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 5,
|
||||
NodeID: 0,
|
||||
Enabled: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
},
|
||||
{
|
||||
Name: "TCP Monitor 2",
|
||||
Type: "tcp",
|
||||
Target: tcpAddr,
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 5,
|
||||
NodeID: 0,
|
||||
Enabled: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
},
|
||||
{
|
||||
Name: "Disabled Monitor",
|
||||
Type: "tcp",
|
||||
Target: "127.0.0.1:1",
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 5,
|
||||
NodeID: 0,
|
||||
Enabled: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
},
|
||||
}
|
||||
|
||||
for _, m := range monitors {
|
||||
if err := r.CreateServiceMonitor(m); err != nil {
|
||||
t.Fatalf("create monitor: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
monitors[2].Enabled = 0
|
||||
if err := r.UpdateServiceMonitor(monitors[2]); err != nil {
|
||||
t.Fatalf("update disabled monitor: %v", err)
|
||||
}
|
||||
|
||||
enabledMonitors, err := r.ListEnabledServiceMonitors()
|
||||
if err != nil {
|
||||
t.Fatalf("list enabled monitors: %v", err)
|
||||
}
|
||||
if len(enabledMonitors) != 2 {
|
||||
t.Fatalf("expected 2 enabled monitors, got %d", len(enabledMonitors))
|
||||
}
|
||||
|
||||
checker := NewChecker(r, nil)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
go checker.Start(ctx)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
results, err := r.GetServiceMonitorResults(monitors[0].ID, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("get tcp results: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatalf("expected at least one result for tcp monitor")
|
||||
}
|
||||
for _, res := range results {
|
||||
if res.Success != 1 {
|
||||
t.Fatalf("expected success for tcp monitor, got failure: %s", res.ErrorMessage)
|
||||
}
|
||||
}
|
||||
|
||||
results2, err := r.GetServiceMonitorResults(monitors[1].ID, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("get tcp results 2: %v", err)
|
||||
}
|
||||
if len(results2) == 0 {
|
||||
t.Fatalf("expected at least one result for tcp monitor 2")
|
||||
}
|
||||
for _, res := range results2 {
|
||||
if res.Success != 1 {
|
||||
t.Fatalf("expected success for tcp monitor 2, got failure: %s", res.ErrorMessage)
|
||||
}
|
||||
}
|
||||
|
||||
disabledResults, err := r.GetServiceMonitorResults(monitors[2].ID, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("get disabled results: %v", err)
|
||||
}
|
||||
if len(disabledResults) != 0 {
|
||||
t.Fatalf("expected no results for disabled monitor, got %d", len(disabledResults))
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckerUnsupportedType(t *testing.T) {
|
||||
checker := NewChecker(nil, nil)
|
||||
limits := checker.loadServiceMonitorLimits()
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
Type: "http",
|
||||
Target: "https://example.com",
|
||||
TimeoutSec: 5,
|
||||
}
|
||||
result := checker.executeCheck(monitor, now, limits)
|
||||
if result.Success == 1 {
|
||||
t.Fatalf("expected failure for unsupported type")
|
||||
}
|
||||
if result.ErrorMessage == "" {
|
||||
t.Fatalf("expected error message for unsupported type")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckerDefaultTimeout(t *testing.T) {
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
addr := listener.Addr().String()
|
||||
|
||||
go func() {
|
||||
for {
|
||||
conn, err := listener.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
conn.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
checker := NewChecker(nil, nil)
|
||||
limits := checker.loadServiceMonitorLimits()
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
Type: "tcp",
|
||||
Target: addr,
|
||||
TimeoutSec: 0,
|
||||
}
|
||||
result := checker.executeCheck(monitor, now, limits)
|
||||
if result.Success != 1 {
|
||||
t.Fatalf("expected success with default timeout, got error: %s", result.ErrorMessage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckerStop(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
go func() {
|
||||
for {
|
||||
conn, err := listener.Accept()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
conn.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
Name: "Test Monitor",
|
||||
Type: "tcp",
|
||||
Target: listener.Addr().String(),
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 5,
|
||||
NodeID: 0,
|
||||
Enabled: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
if err := r.CreateServiceMonitor(monitor); err != nil {
|
||||
t.Fatalf("create monitor: %v", err)
|
||||
}
|
||||
|
||||
checker := NewChecker(r, nil)
|
||||
|
||||
ctx := context.Background()
|
||||
go checker.Start(ctx)
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
checker.Stop()
|
||||
|
||||
results, err := r.GetServiceMonitorResults(monitor.ID, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("get results: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatalf("expected at least one result before stop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckerRunsOnNodeWhenNodeIDSet(t *testing.T) {
|
||||
fake := &fakeCommander{
|
||||
res: ws.CommandResult{
|
||||
Success: true,
|
||||
Data: map[string]interface{}{
|
||||
"success": false,
|
||||
"latencyMs": float64(12),
|
||||
"errorMessage": "unreachable",
|
||||
},
|
||||
},
|
||||
}
|
||||
checker := NewChecker(nil, fake)
|
||||
limits := checker.loadServiceMonitorLimits()
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
ID: 99,
|
||||
Type: "icmp",
|
||||
Target: "8.8.8.8",
|
||||
TimeoutSec: 2,
|
||||
NodeID: 123,
|
||||
}
|
||||
res := checker.executeCheck(monitor, now, limits)
|
||||
if fake.lastNodeID != 123 {
|
||||
t.Fatalf("expected command to be sent to node 123, got %d", fake.lastNodeID)
|
||||
}
|
||||
if fake.lastType != "ServiceMonitorCheck" {
|
||||
t.Fatalf("expected ServiceMonitorCheck command, got %s", fake.lastType)
|
||||
}
|
||||
if res.Success != 0 {
|
||||
t.Fatalf("expected failed result from node check")
|
||||
}
|
||||
if res.ErrorMessage != "unreachable" {
|
||||
t.Fatalf("expected errorMessage unreachable, got %q", res.ErrorMessage)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckerDoesNotBurstOnRestartWhenRecentResultsExist(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
monitor := &model.ServiceMonitor{
|
||||
Name: "recent-monitor",
|
||||
Type: "tcp",
|
||||
Target: "127.0.0.1:1",
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 1,
|
||||
NodeID: 0,
|
||||
Enabled: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
if err := r.CreateServiceMonitor(monitor); err != nil {
|
||||
t.Fatalf("create monitor: %v", err)
|
||||
}
|
||||
if err := r.InsertServiceMonitorResult(&model.ServiceMonitorResult{
|
||||
MonitorID: monitor.ID,
|
||||
NodeID: 0,
|
||||
Timestamp: now - 10_000,
|
||||
Success: 1,
|
||||
}); err != nil {
|
||||
t.Fatalf("seed recent result: %v", err)
|
||||
}
|
||||
|
||||
checker := NewChecker(r, nil)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
go checker.Start(ctx)
|
||||
// Give the initial scan a chance to run.
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
cancel()
|
||||
checker.Stop()
|
||||
|
||||
results, err := r.GetServiceMonitorResults(monitor.ID, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("get results: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Fatalf("expected no immediate rerun (1 result), got %d", len(results))
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckerConcurrencyPreventsSlowMonitorBlockingOthers(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
// Force worker limit to at least 2 for this test.
|
||||
_ = r.UpsertConfig(monitoring.ConfigServiceMonitorWorkerLimit, "2", now)
|
||||
|
||||
slow := &model.ServiceMonitor{
|
||||
Name: "slow",
|
||||
Type: "icmp",
|
||||
Target: "8.8.8.8",
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 1,
|
||||
NodeID: 123,
|
||||
Enabled: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
if err := r.CreateServiceMonitor(slow); err != nil {
|
||||
t.Fatalf("create slow monitor: %v", err)
|
||||
}
|
||||
fast := &model.ServiceMonitor{
|
||||
Name: "fast",
|
||||
Type: "icmp",
|
||||
Target: "1.1.1.1",
|
||||
IntervalSec: 60,
|
||||
TimeoutSec: 1,
|
||||
NodeID: 123,
|
||||
Enabled: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
if err := r.CreateServiceMonitor(fast); err != nil {
|
||||
t.Fatalf("create fast monitor: %v", err)
|
||||
}
|
||||
|
||||
cmd := &delayedCommander{delayByMonitorID: map[int64]time.Duration{slow.ID: 800 * time.Millisecond}}
|
||||
checker := NewChecker(r, cmd)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
go checker.Start(ctx)
|
||||
|
||||
// Fast monitor should complete even while slow one is still running.
|
||||
time.Sleep(250 * time.Millisecond)
|
||||
results, err := r.GetServiceMonitorResults(fast.ID, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("get fast results: %v", err)
|
||||
}
|
||||
if len(results) == 0 {
|
||||
t.Fatalf("expected fast monitor to have results without waiting for slow")
|
||||
}
|
||||
|
||||
cancel()
|
||||
checker.Stop()
|
||||
}
|
||||
@@ -1,10 +1,13 @@
|
||||
# BACKEND HTTP HANDLER KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
HTTP request handlers for FLVX Admin API. Core business logic layer.
|
||||
**Stack:** Go 1.23, net/http, GORM via Repository pattern.
|
||||
**Stack:** Go 1.24, net/http, GORM via Repository pattern.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
@@ -14,7 +17,7 @@ handler/
|
||||
├── federation.go # Federation/cluster sync API
|
||||
├── flow_policy.go # Traffic policy API
|
||||
├── jobs.go # Background job management (sync, cleanup)
|
||||
├── mutations.go # CRUD for users, tunnels, forwards (largest: 100k+ LOC)
|
||||
├── mutations.go # CRUD for users, tunnels, forwards (~3700 LOC)
|
||||
└── upgrade.go # System upgrade API
|
||||
```
|
||||
|
||||
@@ -26,10 +29,11 @@ handler/
|
||||
| **Federation Sync** | `federation.go` | Panel-to-panel sync |
|
||||
| **Traffic Policies** | `flow_policy.go` | Flow limiting, quota management |
|
||||
| **Background Jobs** | `jobs.go` | Scheduled sync/cleanup tasks |
|
||||
| **Node Control** | `control_plane.go` | Node add/delete/list operations |
|
||||
|
||||
## CONVENTIONS
|
||||
- Inherits from parent: GORM via Repository pattern, JWT in Authorization header.
|
||||
- Large files expected (`mutations.go` 3716 LOC - central mutation hub).
|
||||
- Large files expected (`mutations.go` ~3700 LOC - central mutation hub).
|
||||
- Uses `repo.Repository` for DB access via `h.repo.XXX()` methods.
|
||||
- Handlers never call `repo.DB()` directly — all queries go through Repository methods.
|
||||
- Domain-driven file split: one file per functional area (federation, jobs, etc.).
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,11 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildForwardControlServiceNamesPauseResume(t *testing.T) {
|
||||
@@ -42,6 +45,20 @@ func TestBuildForwardServiceBaseCandidatesWithZeroPreferred(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceBaseWithResolvedUserTunnel(t *testing.T) {
|
||||
got := buildForwardServiceBaseWithResolvedUserTunnel(12, 34, 56)
|
||||
if got != "12_34_56" {
|
||||
t.Fatalf("expected 12_34_56, got %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceBaseWithResolvedUserTunnelFallbackToZero(t *testing.T) {
|
||||
got := buildForwardServiceBaseWithResolvedUserTunnel(12, 34, 0)
|
||||
if got != "12_34_0" {
|
||||
t.Fatalf("expected 12_34_0, got %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldTryLegacySingleService(t *testing.T) {
|
||||
if !shouldTryLegacySingleService("PauseService") {
|
||||
t.Fatalf("PauseService should require legacy fallback")
|
||||
@@ -53,3 +70,475 @@ func TestShouldTryLegacySingleService(t *testing.T) {
|
||||
t.Fatalf("DeleteService should not require legacy fallback")
|
||||
}
|
||||
}
|
||||
|
||||
func TestShouldSelfHealForwardServiceControl(t *testing.T) {
|
||||
if !shouldSelfHealForwardServiceControl("PauseService") {
|
||||
t.Fatalf("PauseService should trigger self-heal")
|
||||
}
|
||||
if !shouldSelfHealForwardServiceControl(" resumeService ") {
|
||||
t.Fatalf("ResumeService should trigger self-heal")
|
||||
}
|
||||
if shouldSelfHealForwardServiceControl("DeleteService") {
|
||||
t.Fatalf("DeleteService should not trigger self-heal")
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlForwardServiceCommandHandledOnKnownVariant(t *testing.T) {
|
||||
bases := []string{"12_34_56"}
|
||||
called := make([]string, 0)
|
||||
handled, lastNotFoundErr, err := controlForwardServiceCommand(bases, "PauseService", func(name string) error {
|
||||
called = append(called, name)
|
||||
if name == "12_34_56_udp" {
|
||||
return nil
|
||||
}
|
||||
return errors.New("service " + name + " not found")
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !handled {
|
||||
t.Fatalf("expected handled=true")
|
||||
}
|
||||
if lastNotFoundErr != nil {
|
||||
t.Fatalf("expected lastNotFoundErr=nil when handled")
|
||||
}
|
||||
wantCalls := []string{"12_34_56_tcp", "12_34_56_udp", "12_34_56"}
|
||||
if !reflect.DeepEqual(called, wantCalls) {
|
||||
t.Fatalf("expected calls %v, got %v", wantCalls, called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlForwardServiceCommandReturnsLastNotFoundWhenAllMissing(t *testing.T) {
|
||||
bases := []string{"12_34_56"}
|
||||
handled, lastNotFoundErr, err := controlForwardServiceCommand(bases, "PauseService", func(name string) error {
|
||||
return errors.New("service " + name + " not found")
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if handled {
|
||||
t.Fatalf("expected handled=false")
|
||||
}
|
||||
if lastNotFoundErr == nil {
|
||||
t.Fatalf("expected lastNotFoundErr when all variants are missing")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardServiceCandidatesSkipsNotFoundUntilLegacyMatch(t *testing.T) {
|
||||
bases := []string{"12_34_56", "12_34_0"}
|
||||
called := make([]string, 0)
|
||||
err := deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
called = append(called, name)
|
||||
if name == "12_34_0" {
|
||||
return nil
|
||||
}
|
||||
return errors.New("service " + name + " not found")
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
wantCalls := []string{"12_34_56_tcp", "12_34_56_udp", "12_34_56", "12_34_0_tcp", "12_34_0_udp", "12_34_0"}
|
||||
if !reflect.DeepEqual(called, wantCalls) {
|
||||
t.Fatalf("expected calls %v, got %v", wantCalls, called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardServiceCandidatesTreatsAllMissingAsSuccess(t *testing.T) {
|
||||
bases := []string{"12_34_56", "12_34_0"}
|
||||
err := deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
return errors.New("service " + name + " not found")
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("all-missing delete should be tolerated, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardServiceBaseCandidatesIncludesResolvedAndLegacyZero(t *testing.T) {
|
||||
bases := buildForwardServiceBaseCandidates(46, 9, 123, []int64{123, 77, 0})
|
||||
want := []string{"46_9_123", "46_9_77", "46_9_0"}
|
||||
if !reflect.DeepEqual(bases, want) {
|
||||
t.Fatalf("expected %v, got %v", want, bases)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardServiceBasesOnNodeRetriesLegacyZeroResidue(t *testing.T) {
|
||||
bases := []string{"46_9_123", "46_9_0"}
|
||||
called := make([]string, 0)
|
||||
err := deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
called = append(called, name)
|
||||
if name == "46_9_0_tcp" || name == "46_9_0_udp" {
|
||||
return nil
|
||||
}
|
||||
return errors.New("service " + name + " not found")
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
want := []string{"46_9_123_tcp", "46_9_123_udp", "46_9_123", "46_9_0_tcp", "46_9_0_udp", "46_9_0"}
|
||||
if !reflect.DeepEqual(called, want) {
|
||||
t.Fatalf("expected calls %v, got %v", want, called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardServiceCandidatesDeletesAllMatchingVariants(t *testing.T) {
|
||||
bases := []string{"57_7_7", "57_7_0"}
|
||||
called := make([]string, 0)
|
||||
err := deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
called = append(called, name)
|
||||
switch name {
|
||||
case "57_7_7_tcp", "57_7_7_udp", "57_7_0_tcp", "57_7_0_udp":
|
||||
return nil
|
||||
default:
|
||||
return errors.New("service " + name + " not found")
|
||||
}
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
want := []string{"57_7_7_tcp", "57_7_7_udp", "57_7_7", "57_7_0_tcp", "57_7_0_udp", "57_7_0"}
|
||||
if !reflect.DeepEqual(called, want) {
|
||||
t.Fatalf("expected calls %v, got %v", want, called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateForwardPortAvailabilityRejectsOtherForwardOccupancy(t *testing.T) {
|
||||
h := &Handler{repo: nil}
|
||||
node := &nodeRecord{ID: 9, Name: "test-node"}
|
||||
_ = h
|
||||
_ = node
|
||||
|
||||
rawRepo, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
h = &Handler{repo: rawRepo}
|
||||
if err := rawRepo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(1, 9, 2000)`).Error; err != nil {
|
||||
t.Fatalf("insert forward port: %v", err)
|
||||
}
|
||||
|
||||
err = h.validateForwardPortAvailability(&nodeRecord{ID: 9, Name: "test-node"}, 2000, 2)
|
||||
if err == nil {
|
||||
t.Fatalf("expected occupancy error")
|
||||
}
|
||||
if err.Error() != "节点 test-node 端口 2000 已被其他转发占用" {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
err = h.validateForwardPortAvailability(&nodeRecord{ID: 9, Name: "test-node"}, 2000, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("same forward should be allowed, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestControlForwardServiceCommandReturnsHardError(t *testing.T) {
|
||||
bases := []string{"12_34_56"}
|
||||
handled, lastNotFoundErr, err := controlForwardServiceCommand(bases, "PauseService", func(name string) error {
|
||||
if name == "12_34_56_tcp" {
|
||||
return errors.New("network timeout")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatalf("expected hard error")
|
||||
}
|
||||
if handled {
|
||||
t.Fatalf("expected handled=false on hard error")
|
||||
}
|
||||
if lastNotFoundErr != nil {
|
||||
t.Fatalf("did not expect not-found error alongside hard error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsAlreadyExistsMessage(t *testing.T) {
|
||||
if !isAlreadyExistsMessage("service demo already exists") {
|
||||
t.Fatalf("expected already exists message to be tolerated")
|
||||
}
|
||||
if !isAlreadyExistsMessage("服务已存在") {
|
||||
t.Fatalf("expected Chinese already exists message to be tolerated")
|
||||
}
|
||||
if !isAlreadyExistsMessage("service demo alreadyexists") {
|
||||
t.Fatalf("missing-space alreadyexists should be tolerated")
|
||||
}
|
||||
if isAlreadyExistsMessage("listen tcp [::]:10001: bind: address already in use") {
|
||||
t.Fatalf("address already in use must not be treated as already exists")
|
||||
}
|
||||
if isAlreadyExistsMessage("create service 57_7_7_tcp failed: listen tcp4 0.0.0.0:46222: bind: address alreadyin use") {
|
||||
t.Fatalf("alreadyin-use variant must not be treated as already exists")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBindAddressInUseError(t *testing.T) {
|
||||
if !isBindAddressInUseError(errors.New("listen tcp [::]:10001: bind: address already in use")) {
|
||||
t.Fatalf("address already in use should be detected")
|
||||
}
|
||||
if !isBindAddressInUseError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannot assign requested address")) {
|
||||
t.Fatalf("cannot assign requested address should be detected")
|
||||
}
|
||||
if isBindAddressInUseError(errors.New("service demo already exists")) {
|
||||
t.Fatalf("already exists should not be treated as bind conflict")
|
||||
}
|
||||
if isBindAddressInUseError(nil) {
|
||||
t.Fatalf("nil error should not be treated as bind conflict")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsAddressAlreadyInUseError(t *testing.T) {
|
||||
if !isAddressAlreadyInUseError(errors.New("listen tcp [::]:10001: bind: address already in use")) {
|
||||
t.Fatalf("address already in use should be detected")
|
||||
}
|
||||
if !isAddressAlreadyInUseError(errors.New("create service 57_7_7_tcp failed: listen tcp4 0.0.0.0:46222: bind: address alreadyin use")) {
|
||||
t.Fatalf("missing-space alreadyin-use variant should be detected")
|
||||
}
|
||||
if isAddressAlreadyInUseError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannot assign requested address")) {
|
||||
t.Fatalf("cannot assign requested address should not be treated as address-in-use")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsCannotAssignRequestedAddressError(t *testing.T) {
|
||||
if !isCannotAssignRequestedAddressError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannot assign requested address")) {
|
||||
t.Fatalf("cannot assign requested address should be detected")
|
||||
}
|
||||
if !isCannotAssignRequestedAddressError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannotassignrequestedaddress")) {
|
||||
t.Fatalf("missing-space cannotassignrequestedaddress variant should be detected")
|
||||
}
|
||||
if isCannotAssignRequestedAddressError(errors.New("listen tcp [::]:10001: bind: address already in use")) {
|
||||
t.Fatalf("address already in use should not be treated as cannot-assign")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryTunnelServiceAddWithCleanupRetriesOnAddressInUse(t *testing.T) {
|
||||
addCalls := 0
|
||||
cleanupCalls := 0
|
||||
err := retryTunnelServiceAddWithCleanup(
|
||||
func() error {
|
||||
addCalls++
|
||||
if addCalls == 1 {
|
||||
return errors.New("listen tcp 10.0.0.1:32000: bind: address already in use")
|
||||
}
|
||||
return nil
|
||||
},
|
||||
func() error {
|
||||
cleanupCalls++
|
||||
return nil
|
||||
},
|
||||
0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("expected retry to succeed, got %v", err)
|
||||
}
|
||||
if addCalls != 2 {
|
||||
t.Fatalf("expected 2 add attempts, got %d", addCalls)
|
||||
}
|
||||
if cleanupCalls != 1 {
|
||||
t.Fatalf("expected 1 cleanup attempt, got %d", cleanupCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryTunnelServiceAddWithCleanupSkipsCleanupOnNonBindError(t *testing.T) {
|
||||
addCalls := 0
|
||||
cleanupCalls := 0
|
||||
err := retryTunnelServiceAddWithCleanup(
|
||||
func() error {
|
||||
addCalls++
|
||||
return errors.New("network timeout")
|
||||
},
|
||||
func() error {
|
||||
cleanupCalls++
|
||||
return nil
|
||||
},
|
||||
0,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatalf("expected hard error")
|
||||
}
|
||||
if addCalls != 1 {
|
||||
t.Fatalf("expected 1 add attempt, got %d", addCalls)
|
||||
}
|
||||
if cleanupCalls != 0 {
|
||||
t.Fatalf("expected 0 cleanup attempts, got %d", cleanupCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) {
|
||||
cleanupErr := errors.New("delete failed")
|
||||
err := retryTunnelServiceAddWithCleanup(
|
||||
func() error {
|
||||
return errors.New("listen tcp 10.0.0.1:32000: bind: address already in use")
|
||||
},
|
||||
func() error {
|
||||
return cleanupErr
|
||||
},
|
||||
0,
|
||||
)
|
||||
if !errors.Is(err, cleanupErr) {
|
||||
t.Fatalf("expected cleanup error %v, got %v", cleanupErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, false)
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
for _, svc := range services {
|
||||
addr, _ := svc["addr"].(string)
|
||||
if addr != "10.9.8.7:22000" {
|
||||
t.Fatalf("expected bind IP address 10.9.8.7:22000, got %q", addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, false)
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
tcpAddr, _ := services[0]["addr"].(string)
|
||||
udpAddr, _ := services[1]["addr"].(string)
|
||||
if tcpAddr != "0.0.0.0:22001" {
|
||||
t.Fatalf("expected tcp addr 0.0.0.0:22001, got %q", tcpAddr)
|
||||
}
|
||||
if udpAddr != "[::]:22001" {
|
||||
t.Fatalf("expected udp addr [::]:22001, got %q", udpAddr)
|
||||
}
|
||||
}
|
||||
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil, false)
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
for _, svc := range services {
|
||||
addr, _ := svc["addr"].(string)
|
||||
if addr != "3.3.3.3:12345" {
|
||||
t.Fatalf("expected bind IP with port 3.3.3.3:12345, got %q", addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
bindIP string
|
||||
port int
|
||||
wantAddr string
|
||||
}{
|
||||
{
|
||||
name: "pure ipv6 without port",
|
||||
bindIP: "2001:db8::1",
|
||||
port: 22000,
|
||||
wantAddr: "[2001:db8::1]:22000",
|
||||
},
|
||||
{
|
||||
name: "bracketed ipv6 without port",
|
||||
bindIP: "[2001:db8::2]",
|
||||
port: 22001,
|
||||
wantAddr: "[2001:db8::2]:22001",
|
||||
},
|
||||
{
|
||||
name: "bracketed ipv6 with port",
|
||||
bindIP: "[2001:db8::3]:8080",
|
||||
port: 55555,
|
||||
wantAddr: "[2001:db8::3]:8080",
|
||||
},
|
||||
{
|
||||
name: "ipv6 link-local with zone",
|
||||
bindIP: "fe80::1%eth0",
|
||||
port: 22002,
|
||||
wantAddr: "[fe80::1%eth0]:22002",
|
||||
},
|
||||
{
|
||||
name: "ipv6 localhost",
|
||||
bindIP: "::1",
|
||||
port: 22003,
|
||||
wantAddr: "[::1]:22003",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil, false)
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
for _, svc := range services {
|
||||
addr, _ := svc["addr"].(string)
|
||||
if addr != tt.wantAddr {
|
||||
t.Fatalf("expected addr %q, got %q", tt.wantAddr, addr)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "https with path",
|
||||
in: "https://panel.example.com:8443/api/v1",
|
||||
want: "panel.example.com:8443",
|
||||
},
|
||||
{
|
||||
name: "wss with query",
|
||||
in: "wss://panel.example.com:443/system-info?x=1",
|
||||
want: "panel.example.com:443",
|
||||
},
|
||||
{
|
||||
name: "http without port",
|
||||
in: "http://panel.example.com",
|
||||
want: "panel.example.com",
|
||||
},
|
||||
{
|
||||
name: "manual host with trailing path",
|
||||
in: "panel.example.com:8080/path",
|
||||
want: "panel.example.com:8080",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := processServerAddress(tt.in); got != tt.want {
|
||||
t.Fatalf("%s: expected %q, got %q", tt.name, tt.want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessServerAddress_NormalizesIPv6(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "ipv6 host only",
|
||||
in: "2001:db8::1",
|
||||
want: "[2001:db8::1]",
|
||||
},
|
||||
{
|
||||
name: "ipv6 host and port",
|
||||
in: "https://[2001:db8::1]:8443/path",
|
||||
want: "[2001:db8::1]:8443",
|
||||
},
|
||||
{
|
||||
name: "already bracketed",
|
||||
in: "[2001:db8::2]:9000",
|
||||
want: "[2001:db8::2]:9000",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
if got := processServerAddress(tt.in); got != tt.want {
|
||||
t.Fatalf("%s: expected %q, got %q", tt.name, tt.want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
type diagnosisStreamEvent struct {
|
||||
Type string `json:"type"`
|
||||
Data interface{} `json:"data,omitempty"`
|
||||
TS int64 `json:"ts"`
|
||||
}
|
||||
|
||||
func prepareDiagnosisStreamResponse(w http.ResponseWriter) (http.Flusher, error) {
|
||||
flusher, ok := w.(http.Flusher)
|
||||
if !ok {
|
||||
return nil, errors.New("当前服务不支持流式响应")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/x-ndjson; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
w.Header().Set("X-Accel-Buffering", "no")
|
||||
return flusher, nil
|
||||
}
|
||||
|
||||
func writeDiagnosisStreamEvent(encoder *json.Encoder, flusher http.Flusher, eventType string, data interface{}) error {
|
||||
if encoder == nil || flusher == nil {
|
||||
return errors.New("流式响应写入器未初始化")
|
||||
}
|
||||
event := diagnosisStreamEvent{Type: eventType, Data: data, TS: time.Now().UnixMilli()}
|
||||
if err := encoder.Encode(event); err != nil {
|
||||
return err
|
||||
}
|
||||
flusher.Flush()
|
||||
return nil
|
||||
}
|
||||
|
||||
func summarizeDiagnosisProgress(results []map[string]interface{}) diagnosisProgress {
|
||||
progress := diagnosisProgress{Total: len(results)}
|
||||
for _, item := range results {
|
||||
progress.Completed++
|
||||
if asBool(item["success"], false) {
|
||||
progress.Success++
|
||||
} else {
|
||||
progress.Failed++
|
||||
}
|
||||
}
|
||||
return progress
|
||||
}
|
||||
|
||||
func shouldIgnoreDiagnosisStreamError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
if errors.Is(err, context.Canceled) {
|
||||
return true
|
||||
}
|
||||
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||
if strings.Contains(msg, "broken pipe") || strings.Contains(msg, "connection reset by peer") {
|
||||
return true
|
||||
}
|
||||
if strings.Contains(msg, "stream already closed") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Handler) streamDiagnosisRuntime(ctx context.Context, cancel context.CancelFunc, w http.ResponseWriter, startPayload map[string]interface{}, workItems []diagnosisWorkItem) error {
|
||||
flusher, err := prepareDiagnosisStreamResponse(w)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
encoder := json.NewEncoder(w)
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"total": len(workItems),
|
||||
"timestamp": time.Now().UnixMilli(),
|
||||
"items": h.buildDiagnosisStreamStartItems(workItems),
|
||||
}
|
||||
for key, value := range startPayload {
|
||||
payload[key] = value
|
||||
}
|
||||
if err := writeDiagnosisStreamEvent(encoder, flusher, "start", payload); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
streamBroken := false
|
||||
emitter := func(index int, item map[string]interface{}, progress diagnosisProgress) {
|
||||
if streamBroken {
|
||||
return
|
||||
}
|
||||
itemPayload := map[string]interface{}{
|
||||
"index": index,
|
||||
"result": item,
|
||||
"progress": progress,
|
||||
}
|
||||
if err := writeDiagnosisStreamEvent(encoder, flusher, "item", itemPayload); err != nil {
|
||||
streamBroken = true
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
results := h.runDiagnosisWorkItems(ctx, workItems, emitter)
|
||||
if streamBroken {
|
||||
return context.Canceled
|
||||
}
|
||||
|
||||
progress := summarizeDiagnosisProgress(results)
|
||||
donePayload := map[string]interface{}{
|
||||
"progress": progress,
|
||||
"timedOut": errors.Is(ctx.Err(), context.DeadlineExceeded),
|
||||
}
|
||||
return writeDiagnosisStreamEvent(encoder, flusher, "done", donePayload)
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelDiagnoseStream(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
id := asInt64FromBodyKey(r, w, "tunnelId")
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
tunnelName, tunnelType, workItems, err := h.prepareTunnelDiagnosis(id)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不完整") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(r.Context(), diagnosisRequestTimeout)
|
||||
defer cancel()
|
||||
|
||||
startPayload := map[string]interface{}{
|
||||
"tunnelName": tunnelName,
|
||||
"tunnelType": tunnelType,
|
||||
}
|
||||
if err := h.streamDiagnosisRuntime(ctx, cancel, w, startPayload, workItems); err != nil {
|
||||
if shouldIgnoreDiagnosisStreamError(err) {
|
||||
return
|
||||
}
|
||||
if strings.Contains(err.Error(), "不支持流式响应") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) forwardDiagnoseStream(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
id := asInt64FromBodyKey(r, w, "forwardId")
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
forward, _, _, err := h.resolveForwardAccess(r, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "不存在") || strings.Contains(err.Error(), "不能为空") || strings.Contains(err.Error(), "错误") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(r.Context(), diagnosisRequestTimeout)
|
||||
defer cancel()
|
||||
|
||||
startPayload := map[string]interface{}{
|
||||
"forwardName": forwardName,
|
||||
}
|
||||
if err := h.streamDiagnosisRuntime(ctx, cancel, w, startPayload, workItems); err != nil {
|
||||
if shouldIgnoreDiagnosisStreamError(err) {
|
||||
return
|
||||
}
|
||||
if strings.Contains(err.Error(), "不支持流式响应") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -8,9 +8,100 @@ import (
|
||||
// nodeSupportsV4 / nodeSupportsV6
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestNodeSupportsV4_Nil(t *testing.T) {
|
||||
if nodeSupportsV4(nil) {
|
||||
t.Fatal("nil node must not support v4")
|
||||
func TestSelectTunnelDialHost_ConnectIpPriority(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// Empty connectIp should be ignored, IP preference takes effect
|
||||
host, err := selectTunnelDialHost(from, to, "", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("empty connectIp should be ignored (v4 preference applies), got %q", host)
|
||||
}
|
||||
// Non-empty connectIp should override IP preference
|
||||
host, err = selectTunnelDialHost(from, to, "v6", "192.168.0.3")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if host != "192.168.0.3" {
|
||||
t.Fatalf("connectIp should override v6 preference, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainServiceConfig_UsesConnectIPForListen(t *testing.T) {
|
||||
node := &nodeRecord{TCPListenAddr: "[::]"}
|
||||
chain := tunnelRuntimeNode{Protocol: "tls", Port: 21000, ConnectIP: "2001:db8::88"}
|
||||
services := buildTunnelChainServiceConfig(99, chain, node, 1)
|
||||
if len(services) != 1 {
|
||||
t.Fatalf("expected 1 service, got %d", len(services))
|
||||
}
|
||||
addr, _ := services[0]["addr"].(string)
|
||||
if addr != "[2001:db8::88]:21000" {
|
||||
t.Fatalf("expected connectIp listen [2001:db8::88]:21000, got %q", addr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainServiceConfig_FallsBackToNodeListenAddr(t *testing.T) {
|
||||
node := &nodeRecord{TCPListenAddr: "10.8.0.5"}
|
||||
chain := tunnelRuntimeNode{Protocol: "tls", Port: 21002}
|
||||
services := buildTunnelChainServiceConfig(99, chain, node, 1)
|
||||
if len(services) != 1 {
|
||||
t.Fatalf("expected 1 service, got %d", len(services))
|
||||
}
|
||||
addr, _ := services[0]["addr"].(string)
|
||||
if addr != "10.8.0.5:21002" {
|
||||
t.Fatalf("expected node listen addr 10.8.0.5:21002, got %q", addr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainServiceConfig_DefaultListenAddrWhenConnectIPEmpty(t *testing.T) {
|
||||
node := &nodeRecord{TCPListenAddr: "[::]"}
|
||||
chain := tunnelRuntimeNode{Protocol: "tls", Port: 21001}
|
||||
services := buildTunnelChainServiceConfig(99, chain, node, 1)
|
||||
if len(services) != 1 {
|
||||
t.Fatalf("expected 1 service, got %d", len(services))
|
||||
}
|
||||
addr, _ := services[0]["addr"].(string)
|
||||
if addr != "[::]:21001" {
|
||||
t.Fatalf("expected default listen [::]:21001, got %q", addr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainServiceConfig_SetsRetriesWhenMultipleCandidates(t *testing.T) {
|
||||
node := &nodeRecord{TCPListenAddr: "[::]"}
|
||||
chain := tunnelRuntimeNode{Protocol: "tls", Port: 21001}
|
||||
services := buildTunnelChainServiceConfig(99, chain, node, 3)
|
||||
if len(services) != 1 {
|
||||
t.Fatalf("expected 1 service, got %d", len(services))
|
||||
}
|
||||
handler, _ := services[0]["handler"].(map[string]interface{})
|
||||
if handler == nil {
|
||||
t.Fatal("expected handler config")
|
||||
}
|
||||
retries, ok := handler["retries"].(int)
|
||||
if !ok {
|
||||
t.Fatal("expected retries to be set when nextHopCandidateCount > 1")
|
||||
}
|
||||
if retries != 2 {
|
||||
t.Fatalf("expected retries=2 (candidates-1), got %d", retries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainServiceConfig_NoRetriesWhenSingleCandidate(t *testing.T) {
|
||||
node := &nodeRecord{TCPListenAddr: "[::]"}
|
||||
chain := tunnelRuntimeNode{Protocol: "tls", Port: 21001}
|
||||
services := buildTunnelChainServiceConfig(99, chain, node, 1)
|
||||
if len(services) != 1 {
|
||||
t.Fatalf("expected 1 service, got %d", len(services))
|
||||
}
|
||||
handler, _ := services[0]["handler"].(map[string]interface{})
|
||||
if handler == nil {
|
||||
t.Fatal("expected handler config")
|
||||
}
|
||||
if _, hasRetries := handler["retries"]; hasRetries {
|
||||
t.Fatal("expected no retries when nextHopCandidateCount is 1")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -23,14 +114,14 @@ func TestNodeSupportsV6_Nil(t *testing.T) {
|
||||
func TestNodeSupportsV4_ExplicitV4(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv4: "10.0.0.1"}
|
||||
if !nodeSupportsV4(n) {
|
||||
t.Fatal("explicit server_ip_v4 must support v4")
|
||||
t.Fatal("explicit server_ip_v4 needs support v4")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeSupportsV6_ExplicitV6(t *testing.T) {
|
||||
n := &nodeRecord{ServerIPv6: "2001:db8::1"}
|
||||
if !nodeSupportsV6(n) {
|
||||
t.Fatal("explicit server_ip_v6 must support v6")
|
||||
t.Fatal("explicit server_ip_v6 needs support v6")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -68,7 +159,7 @@ func TestNodeSupportsV4_LegacyV4Only(t *testing.T) {
|
||||
t.Fatal("legacy v4 ip in server_ip must support v4")
|
||||
}
|
||||
if nodeSupportsV6(n) {
|
||||
t.Fatal("legacy v4 ip in server_ip must not support v6")
|
||||
t.Fatal("legacy v4 ip in server_ip should not support v6")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,7 +169,7 @@ func TestNodeSupportsV6_LegacyV6Only(t *testing.T) {
|
||||
t.Fatal("legacy v6 ip in server_ip must support v6")
|
||||
}
|
||||
if nodeSupportsV4(n) {
|
||||
t.Fatal("legacy v6 ip in server_ip must not support v4")
|
||||
t.Fatal("legacy v6 ip in server_ip should not support v4")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -177,15 +268,15 @@ func v6OnlyNode(name, v6 string) *nodeRecord {
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_NilNodes(t *testing.T) {
|
||||
_, err := selectTunnelDialHost(nil, nil, "")
|
||||
_, err := selectTunnelDialHost(nil, nil, "", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil nodes")
|
||||
}
|
||||
_, err = selectTunnelDialHost(dualStackNode("a", "1.1.1.1", "::1"), nil, "")
|
||||
_, err = selectTunnelDialHost(dualStackNode("a", "1.1.1.1", "::1"), nil, "", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil toNode")
|
||||
}
|
||||
_, err = selectTunnelDialHost(nil, dualStackNode("b", "1.1.1.1", "::1"), "")
|
||||
_, err = selectTunnelDialHost(nil, dualStackNode("b", "1.1.1.1", "::1"), "", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil fromNode")
|
||||
}
|
||||
@@ -194,8 +285,7 @@ func TestSelectTunnelDialHost_NilNodes(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_DualStack_DefaultPreference(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
host, err := selectTunnelDialHost(from, to, "")
|
||||
host, err := selectTunnelDialHost(from, to, "", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -208,8 +298,7 @@ func TestSelectTunnelDialHost_DualStack_DefaultPreference(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_DualStack_PreferV4(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
host, err := selectTunnelDialHost(from, to, "v4", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -221,8 +310,7 @@ func TestSelectTunnelDialHost_DualStack_PreferV4(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_DualStack_PreferV6(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
host, err := selectTunnelDialHost(from, to, "v6", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -234,9 +322,8 @@ func TestSelectTunnelDialHost_DualStack_PreferV6(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_V4Only_PreferV6Fallback(t *testing.T) {
|
||||
from := v4OnlyNode("from", "10.0.0.1")
|
||||
to := v4OnlyNode("to", "10.0.0.2")
|
||||
|
||||
// User prefers v6, but both nodes are v4-only — should fallback to v4
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
host, err := selectTunnelDialHost(from, to, "v6", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -248,9 +335,8 @@ func TestSelectTunnelDialHost_V4Only_PreferV6Fallback(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_V6Only_PreferV4Fallback(t *testing.T) {
|
||||
from := v6OnlyNode("from", "2001:db8::1")
|
||||
to := v6OnlyNode("to", "2001:db8::2")
|
||||
|
||||
// User prefers v4, but both nodes are v6-only — should fallback to v6
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
host, err := selectTunnelDialHost(from, to, "v4", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -259,32 +345,47 @@ func TestSelectTunnelDialHost_V6Only_PreferV4Fallback(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_Incompatible(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_CrossVersion_V4ToV6(t *testing.T) {
|
||||
// v4-only -> v6-only: 跨版本支持,应成功返回 v6 地址
|
||||
from := v4OnlyNode("from", "10.0.0.1")
|
||||
to := v6OnlyNode("to", "2001:db8::2")
|
||||
|
||||
_, err := selectTunnelDialHost(from, to, "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for incompatible nodes (v4-only -> v6-only)")
|
||||
host, err := selectTunnelDialHost(from, to, "", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error for cross-version (v4-only -> v6-only): %v", err)
|
||||
}
|
||||
if host != "2001:db8::2" {
|
||||
t.Fatalf("expected v6 address for cross-version, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_Incompatible_Reverse(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_CrossVersion_V6ToV4(t *testing.T) {
|
||||
// v6-only -> v4-only: 跨版本支持,应成功返回 v4 地址
|
||||
from := v6OnlyNode("from", "2001:db8::1")
|
||||
to := v4OnlyNode("to", "10.0.0.2")
|
||||
host, err := selectTunnelDialHost(from, to, "", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error for cross-version (v6-only -> v4-only): %v", err)
|
||||
}
|
||||
if host != "10.0.0.2" {
|
||||
t.Fatalf("expected v4 address for cross-version, got %q", host)
|
||||
}
|
||||
}
|
||||
|
||||
_, err := selectTunnelDialHost(from, to, "")
|
||||
func TestSelectTunnelDialHost_TrulyIncompatible(t *testing.T) {
|
||||
// 真正不兼容:两个节点都没有任何 IP
|
||||
from := &nodeRecord{Name: "empty-from", ServerIPv4: "", ServerIPv6: "", ServerIP: ""}
|
||||
to := &nodeRecord{Name: "empty-to", ServerIPv4: "", ServerIPv6: "", ServerIP: ""}
|
||||
_, err := selectTunnelDialHost(from, to, "", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for incompatible nodes (v6-only -> v4-only)")
|
||||
t.Fatal("expected error for nodes with no IP addresses")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectTunnelDialHost_WhitespacePreference(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// Whitespace should be trimmed, treated as "v6"
|
||||
host, err := selectTunnelDialHost(from, to, " v6 ")
|
||||
host, err := selectTunnelDialHost(from, to, " v6 ", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -296,9 +397,8 @@ func TestSelectTunnelDialHost_WhitespacePreference(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_MixedStack_FromDualToV4(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := v4OnlyNode("to", "10.0.0.2")
|
||||
|
||||
// v6 preferred, but target only has v4 — should succeed with v4
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
host, err := selectTunnelDialHost(from, to, "v6", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -310,9 +410,8 @@ func TestSelectTunnelDialHost_MixedStack_FromDualToV4(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_MixedStack_FromDualToV6(t *testing.T) {
|
||||
from := dualStackNode("from", "10.0.0.1", "2001:db8::1")
|
||||
to := v6OnlyNode("to", "2001:db8::2")
|
||||
|
||||
// v4 preferred, but target only has v6 — should succeed with v6
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
host, err := selectTunnelDialHost(from, to, "v4", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -324,9 +423,8 @@ func TestSelectTunnelDialHost_MixedStack_FromDualToV6(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_MixedStack_FromV4ToDual(t *testing.T) {
|
||||
from := v4OnlyNode("from", "10.0.0.1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// v6 preferred, but from only has v4 — should use v4 (from can only reach v4 of target)
|
||||
host, err := selectTunnelDialHost(from, to, "v6")
|
||||
host, err := selectTunnelDialHost(from, to, "v6", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -338,9 +436,8 @@ func TestSelectTunnelDialHost_MixedStack_FromV4ToDual(t *testing.T) {
|
||||
func TestSelectTunnelDialHost_MixedStack_FromV6ToDual(t *testing.T) {
|
||||
from := v6OnlyNode("from", "2001:db8::1")
|
||||
to := dualStackNode("to", "10.0.0.2", "2001:db8::2")
|
||||
|
||||
// v4 preferred, but from only has v6 — should use v6
|
||||
host, err := selectTunnelDialHost(from, to, "v4")
|
||||
host, err := selectTunnelDialHost(from, to, "v4", "")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
@@ -367,7 +464,6 @@ func TestNodeDisplayName_Named(t *testing.T) {
|
||||
t.Fatalf("expected 'hk-node', got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeDisplayName_Unnamed(t *testing.T) {
|
||||
n := &nodeRecord{ID: 42}
|
||||
got := nodeDisplayName(n)
|
||||
|
||||
@@ -141,6 +141,32 @@ type remoteUsageNodeItem struct {
|
||||
SyncError string `json:"syncError,omitempty"`
|
||||
}
|
||||
|
||||
func buildFederationServiceConfig(serviceName, addr, protocol, role, chainName string, targetCount int, interfaceName string) map[string]interface{} {
|
||||
service := map[string]interface{}{
|
||||
"name": serviceName,
|
||||
"addr": addr,
|
||||
"handler": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
"listener": map[string]interface{}{
|
||||
"type": protocol,
|
||||
},
|
||||
}
|
||||
if isTLSTunnelProtocol(protocol) {
|
||||
service["handler"].(map[string]interface{})["metadata"] = map[string]interface{}{"nodelay": true}
|
||||
}
|
||||
if role == "middle" {
|
||||
service["handler"].(map[string]interface{})["chain"] = chainName
|
||||
if targetCount > 1 {
|
||||
service["handler"].(map[string]interface{})["retries"] = targetCount - 1
|
||||
}
|
||||
}
|
||||
if role == "exit" && strings.TrimSpace(interfaceName) != "" {
|
||||
service["metadata"] = map[string]interface{}{"interface": interfaceName}
|
||||
}
|
||||
return service
|
||||
}
|
||||
|
||||
func (h *Handler) federationShareList(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid method"))
|
||||
@@ -477,9 +503,14 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
forwardPortRows, err := h.repo.ListActiveForwardPortsForNode(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
usedSet := make(map[int]struct{})
|
||||
bindings := make([]remoteUsageBindingItem, 0, len(bindingRows))
|
||||
bindings := make([]remoteUsageBindingItem, 0, len(bindingRows)+len(forwardPortRows))
|
||||
for _, b := range bindingRows {
|
||||
bindings = append(bindings, remoteUsageBindingItem{
|
||||
BindingID: b.ID,
|
||||
@@ -496,6 +527,29 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque
|
||||
usedSet[b.AllocatedPort] = struct{}{}
|
||||
}
|
||||
}
|
||||
for _, fp := range forwardPortRows {
|
||||
bindings = append(bindings, remoteUsageBindingItem{
|
||||
BindingID: -fp.ForwardID,
|
||||
TunnelID: fp.TunnelID,
|
||||
TunnelName: fp.TunnelName,
|
||||
ChainType: 1,
|
||||
HopInx: 0,
|
||||
AllocatedPort: fp.Port,
|
||||
ResourceKey: fmt.Sprintf("forward:%d", fp.ForwardID),
|
||||
RemoteBindingID: "",
|
||||
UpdatedTime: fp.UpdatedTime,
|
||||
})
|
||||
if fp.Port > 0 {
|
||||
usedSet[fp.Port] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
sort.Slice(bindings, func(i, j int) bool {
|
||||
if bindings[i].AllocatedPort == bindings[j].AllocatedPort {
|
||||
return bindings[i].BindingID < bindings[j].BindingID
|
||||
}
|
||||
return bindings[i].AllocatedPort < bindings[j].AllocatedPort
|
||||
})
|
||||
|
||||
usedPorts := make([]int, 0, len(usedSet))
|
||||
for port := range usedSet {
|
||||
@@ -766,6 +820,37 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
usedPorts, err := h.repo.ListUsedPortsOnNode(share.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
for _, port := range usedPorts {
|
||||
if port == req.RemotePort {
|
||||
response.WriteJSON(w, response.Err(403, "Port already in use"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
runtimeOnPort, err := h.repo.GetActiveForwardPeerShareRuntimeByPort(share.ID, req.RemotePort)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if runtimeOnPort != nil {
|
||||
response.WriteJSON(w, response.Err(403, "Port already in use"))
|
||||
return
|
||||
}
|
||||
existsOnNodePort, err := h.repo.ExistsActivePeerShareRuntimeOnNodePort(share.NodeID, req.RemotePort)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if existsOnNodePort {
|
||||
response.WriteJSON(w, response.Err(403, "Port already in use"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
tunnelID, err := h.repo.CreateFederationTunnel(
|
||||
fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort),
|
||||
@@ -780,6 +865,30 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
|
||||
return
|
||||
}
|
||||
|
||||
runtime := &repo.PeerShareRuntime{
|
||||
ShareID: share.ID,
|
||||
NodeID: share.NodeID,
|
||||
ReservationID: randomToken(24),
|
||||
ResourceKey: fmt.Sprintf("federation-forward-%d-%d-%d", share.ID, tunnelID, req.RemotePort),
|
||||
BindingID: "",
|
||||
Role: "forward",
|
||||
ChainName: "",
|
||||
ServiceName: "",
|
||||
Protocol: defaultString(req.Protocol, "tcp"),
|
||||
Strategy: "fifo",
|
||||
Port: req.RemotePort,
|
||||
Target: strings.TrimSpace(req.Target),
|
||||
Applied: 0,
|
||||
Status: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
if err := h.repo.CreatePeerShareRuntime(runtime); err != nil {
|
||||
_ = h.deleteTunnelByID(tunnelID)
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
h.wsServer.SendCommand(share.NodeID, "reload", nil, time.Second*5)
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
@@ -1013,25 +1122,16 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
}
|
||||
}
|
||||
|
||||
service := map[string]interface{}{
|
||||
"name": serviceName,
|
||||
"addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, runtime.Port),
|
||||
"handler": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
"listener": map[string]interface{}{
|
||||
"type": protocol,
|
||||
},
|
||||
}
|
||||
if isTLSTunnelProtocol(protocol) {
|
||||
service["handler"].(map[string]interface{})["metadata"] = map[string]interface{}{"nodelay": true}
|
||||
}
|
||||
if req.Role == "middle" {
|
||||
service["handler"].(map[string]interface{})["chain"] = chainName
|
||||
}
|
||||
if req.Role == "exit" && strings.TrimSpace(node.InterfaceName) != "" {
|
||||
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
|
||||
}
|
||||
targetCount := len(req.Targets)
|
||||
service := buildFederationServiceConfig(
|
||||
serviceName,
|
||||
fmt.Sprintf("%s:%d", node.TCPListenAddr, runtime.Port),
|
||||
protocol,
|
||||
req.Role,
|
||||
chainName,
|
||||
targetCount,
|
||||
node.InterfaceName,
|
||||
)
|
||||
if _, err := h.sendNodeCommand(share.NodeID, "AddService", []map[string]interface{}{service}, true, false); err != nil {
|
||||
if req.Role == "middle" {
|
||||
_, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
|
||||
@@ -1149,16 +1249,20 @@ func (h *Handler) federationRuntimeDiagnose(w http.ResponseWriter, r *http.Reque
|
||||
if req.Count <= 0 {
|
||||
req.Count = 4
|
||||
}
|
||||
if req.Timeout <= 0 {
|
||||
req.Timeout = 5000
|
||||
if req.Timeout <= 0 || req.Timeout > int(diagnosisCommandTimeout/time.Millisecond) {
|
||||
req.Timeout = int(diagnosisCommandTimeout / time.Millisecond)
|
||||
}
|
||||
commandTimeout := time.Duration(req.Timeout) * time.Millisecond
|
||||
if commandTimeout <= 0 || commandTimeout > diagnosisCommandTimeout {
|
||||
commandTimeout = diagnosisCommandTimeout
|
||||
}
|
||||
|
||||
res, err := h.sendNodeCommand(share.NodeID, "TcpPing", map[string]interface{}{
|
||||
res, err := h.sendNodeCommandWithTimeout(share.NodeID, "TcpPing", map[string]interface{}{
|
||||
"ip": req.IP,
|
||||
"port": req.Port,
|
||||
"count": req.Count,
|
||||
"timeout": req.Timeout,
|
||||
}, false, false)
|
||||
}, commandTimeout, false, false)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
@@ -1211,12 +1315,185 @@ func (h *Handler) federationRuntimeCommand(w http.ResponseWriter, r *http.Reques
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if strings.EqualFold(cmd, "addservice") || strings.EqualFold(cmd, "updateservice") {
|
||||
h.bindPeerShareForwardRuntimeServices(share, req.Data)
|
||||
} else if strings.EqualFold(cmd, "deleteservice") {
|
||||
h.releasePeerShareForwardRuntimeServices(share, req.Data)
|
||||
}
|
||||
response.WriteJSON(w, response.OK(res))
|
||||
}
|
||||
|
||||
type federationForwardServiceBinding struct {
|
||||
Name string
|
||||
Port int
|
||||
}
|
||||
|
||||
func extractFederationServiceEntries(data interface{}) []map[string]interface{} {
|
||||
if data == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if entries := asMapSlice(data); len(entries) > 0 {
|
||||
return entries
|
||||
}
|
||||
|
||||
dataMap, ok := data.(map[string]interface{})
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
if entries := asMapSlice(dataMap["services"]); len(entries) > 0 {
|
||||
return entries
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseFederationForwardServiceBindings(data interface{}) []federationForwardServiceBinding {
|
||||
serviceList := extractFederationServiceEntries(data)
|
||||
bindings := make([]federationForwardServiceBinding, 0, len(serviceList))
|
||||
for _, svcMap := range serviceList {
|
||||
name := normalizeForwardRuntimeServiceName(asString(svcMap["name"]))
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
if _, _, _, ok := parseFlowServiceIDs(name); !ok {
|
||||
continue
|
||||
}
|
||||
addr := strings.TrimSpace(asString(svcMap["addr"]))
|
||||
if addr == "" {
|
||||
continue
|
||||
}
|
||||
_, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
port, err := strconv.Atoi(portStr)
|
||||
if err != nil || port <= 0 {
|
||||
continue
|
||||
}
|
||||
bindings = append(bindings, federationForwardServiceBinding{Name: name, Port: port})
|
||||
}
|
||||
return bindings
|
||||
}
|
||||
|
||||
func parseFederationForwardServiceNamesForRelease(data interface{}) []string {
|
||||
names := make(map[string]struct{})
|
||||
appendName := func(raw string) {
|
||||
name := normalizeForwardRuntimeServiceName(raw)
|
||||
if name == "" {
|
||||
return
|
||||
}
|
||||
if _, _, _, ok := parseFlowServiceIDs(name); !ok {
|
||||
return
|
||||
}
|
||||
names[name] = struct{}{}
|
||||
}
|
||||
|
||||
for _, svcMap := range extractFederationServiceEntries(data) {
|
||||
appendName(asString(svcMap["name"]))
|
||||
}
|
||||
|
||||
if dataMap, ok := data.(map[string]interface{}); ok {
|
||||
for _, item := range asAnySlice(dataMap["services"]) {
|
||||
appendName(asString(item))
|
||||
}
|
||||
}
|
||||
|
||||
for _, item := range asAnySlice(data) {
|
||||
appendName(asString(item))
|
||||
}
|
||||
|
||||
if len(names) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
out := make([]string, 0, len(names))
|
||||
for name := range names {
|
||||
out = append(out, name)
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func (h *Handler) bindPeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) {
|
||||
if h == nil || h.repo == nil || share == nil {
|
||||
return
|
||||
}
|
||||
bindings := parseFederationForwardServiceBindings(data)
|
||||
if len(bindings) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
for _, binding := range bindings {
|
||||
runtime, err := h.repo.GetActiveForwardPeerShareRuntimeByPort(share.ID, binding.Port)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if runtime == nil {
|
||||
runtime, err = h.repo.GetActiveForwardPeerShareRuntimeByServiceName(share.ID, binding.Name)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if runtime == nil {
|
||||
_ = h.repo.CreatePeerShareRuntime(&repo.PeerShareRuntime{
|
||||
ShareID: share.ID,
|
||||
NodeID: share.NodeID,
|
||||
ReservationID: randomToken(24),
|
||||
ResourceKey: fmt.Sprintf("forward-runtime:%d:%s:%d:%s", share.ID, binding.Name, binding.Port, randomToken(8)),
|
||||
BindingID: "",
|
||||
Role: "forward",
|
||||
ChainName: "",
|
||||
ServiceName: binding.Name,
|
||||
Protocol: "tcp",
|
||||
Strategy: "fifo",
|
||||
Port: binding.Port,
|
||||
Target: "",
|
||||
Applied: 1,
|
||||
Status: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
})
|
||||
continue
|
||||
}
|
||||
if runtime.ServiceName == binding.Name && runtime.Applied == 1 && runtime.Port == binding.Port && runtime.Status == 1 {
|
||||
continue
|
||||
}
|
||||
runtime.ServiceName = binding.Name
|
||||
runtime.Port = binding.Port
|
||||
runtime.Applied = 1
|
||||
runtime.Status = 1
|
||||
runtime.UpdatedTime = now
|
||||
if strings.TrimSpace(runtime.Protocol) == "" {
|
||||
runtime.Protocol = "tcp"
|
||||
}
|
||||
if strings.TrimSpace(runtime.Strategy) == "" {
|
||||
runtime.Strategy = "fifo"
|
||||
}
|
||||
_ = h.repo.UpdatePeerShareRuntime(runtime)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) releasePeerShareForwardRuntimeServices(share *repo.PeerShare, data interface{}) {
|
||||
if h == nil || h.repo == nil || share == nil {
|
||||
return
|
||||
}
|
||||
names := parseFederationForwardServiceNamesForRelease(data)
|
||||
if len(names) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
for _, name := range names {
|
||||
_ = h.repo.MarkForwardPeerShareRuntimeReleasedByServiceName(share.ID, name, now)
|
||||
}
|
||||
}
|
||||
|
||||
func isFederationRuntimeCommandAllowed(commandType string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(commandType)) {
|
||||
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "deletelimiters", "tcpping", "reload":
|
||||
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice", "addchains", "deletechains", "addlimiters", "updatelimiters", "deletelimiters", "tcpping", "reload":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
@@ -1236,36 +1513,26 @@ func validateFederationCommandPorts(share *repo.PeerShare, data interface{}) err
|
||||
if share == nil || (share.PortRangeStart <= 0 && share.PortRangeEnd <= 0) {
|
||||
return nil
|
||||
}
|
||||
dataMap, ok := data.(map[string]interface{})
|
||||
if !ok {
|
||||
|
||||
serviceList := extractFederationServiceEntries(data)
|
||||
if len(serviceList) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
if services, ok := dataMap["services"]; ok {
|
||||
serviceList, ok := services.([]interface{})
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid services format")
|
||||
for _, svcMap := range serviceList {
|
||||
addr := asString(svcMap["addr"])
|
||||
if addr == "" {
|
||||
continue
|
||||
}
|
||||
for _, svc := range serviceList {
|
||||
svcMap, ok := svc.(map[string]interface{})
|
||||
if !ok {
|
||||
return fmt.Errorf("invalid service entry format")
|
||||
}
|
||||
addr, ok := svcMap["addr"].(string)
|
||||
if !ok || addr == "" {
|
||||
continue
|
||||
}
|
||||
_, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid service address: %s", addr)
|
||||
}
|
||||
port, err := strconv.Atoi(portStr)
|
||||
if err != nil || port <= 0 {
|
||||
return fmt.Errorf("invalid port in service address: %s", addr)
|
||||
}
|
||||
if port < share.PortRangeStart || port > share.PortRangeEnd {
|
||||
return fmt.Errorf("port %d out of allowed range %d-%d", port, share.PortRangeStart, share.PortRangeEnd)
|
||||
}
|
||||
_, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid service address: %s", addr)
|
||||
}
|
||||
port, err := strconv.Atoi(portStr)
|
||||
if err != nil || port <= 0 {
|
||||
return fmt.Errorf("invalid port in service address: %s", addr)
|
||||
}
|
||||
if port < share.PortRangeStart || port > share.PortRangeEnd {
|
||||
return fmt.Errorf("port %d out of allowed range %d-%d", port, share.PortRangeStart, share.PortRangeEnd)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -227,6 +227,60 @@ func TestPrepareTunnelCreateStateAllowsOfflineRemoteMiddleNode(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationServiceConfig_MiddleRoleWithMultipleTargets_SetsRetries(t *testing.T) {
|
||||
service := buildFederationServiceConfig("svc-middle", ":40000", "tls", "middle", "chain-next", 3, "")
|
||||
handler := service["handler"].(map[string]interface{})
|
||||
if handler["chain"] != "chain-next" {
|
||||
t.Fatalf("expected chain 'chain-next', got %v", handler["chain"])
|
||||
}
|
||||
if handler["retries"] != 2 {
|
||||
t.Fatalf("expected retries 2 for 3 targets, got %v", handler["retries"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationServiceConfig_MiddleRoleWithSingleTarget_NoRetries(t *testing.T) {
|
||||
service := buildFederationServiceConfig("svc-middle", ":40000", "tls", "middle", "chain-next", 1, "")
|
||||
handler := service["handler"].(map[string]interface{})
|
||||
if handler["chain"] != "chain-next" {
|
||||
t.Fatalf("expected chain 'chain-next', got %v", handler["chain"])
|
||||
}
|
||||
if _, hasRetries := handler["retries"]; hasRetries {
|
||||
t.Fatalf("expected no retries for single target, got %v", handler["retries"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationServiceConfig_ExitRole_NoRetriesRegardlessOfTargets(t *testing.T) {
|
||||
service := buildFederationServiceConfig("svc-exit", ":40000", "tls", "exit", "", 3, "eth0")
|
||||
handler := service["handler"].(map[string]interface{})
|
||||
if _, hasChain := handler["chain"]; hasChain {
|
||||
t.Fatalf("expected no chain for exit role, got %v", handler["chain"])
|
||||
}
|
||||
if _, hasRetries := handler["retries"]; hasRetries {
|
||||
t.Fatalf("expected no retries for exit role, got %v", handler["retries"])
|
||||
}
|
||||
metadata := service["metadata"].(map[string]interface{})
|
||||
if metadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface 'eth0', got %v", metadata["interface"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationServiceConfig_TLSTunnelProtocol_SetsNodelay(t *testing.T) {
|
||||
service := buildFederationServiceConfig("svc-tls", ":40000", "tls", "middle", "chain-next", 2, "")
|
||||
handler := service["handler"].(map[string]interface{})
|
||||
meta := handler["metadata"].(map[string]interface{})
|
||||
if meta["nodelay"] != true {
|
||||
t.Fatalf("expected nodelay=true for TLS protocol, got %v", meta["nodelay"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationServiceConfig_NonTLSProtocol_NoNodelay(t *testing.T) {
|
||||
service := buildFederationServiceConfig("svc-tcp", ":40000", "tcp", "middle", "chain-next", 2, "")
|
||||
handler := service["handler"].(map[string]interface{})
|
||||
if _, hasMeta := handler["metadata"]; hasMeta {
|
||||
t.Fatalf("expected no metadata for non-TLS protocol, got %v", handler["metadata"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
|
||||
@@ -414,6 +414,445 @@ func TestFederationShareResetFlow(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationTunnelCreateCreatesPeerShareRuntime(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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, is_remote, remote_url, remote_token, remote_config)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "federation-forward-node", "federation-forward-secret", "10.90.80.70", "10.90.80.70", "", "24000-24020", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "federation-forward-node")
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "federation-forward-share",
|
||||
NodeID: nodeID,
|
||||
Token: "federation-forward-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 24000,
|
||||
PortRangeEnd: 24020,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("federation-forward-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
body, err := json.Marshal(federationTunnelRequest{
|
||||
Protocol: "tcp",
|
||||
RemotePort: 24001,
|
||||
Target: "1.1.1.1:443",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/tunnel/create", bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", "Bearer "+share.Token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
h.federationTunnelCreate(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
|
||||
var payload response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected response code 0, got %d (%s)", payload.Code, payload.Msg)
|
||||
}
|
||||
|
||||
runtimeCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND port = ? AND status = 1`, share.ID, 24001)
|
||||
if runtimeCount != 1 {
|
||||
t.Fatalf("expected 1 runtime row for new federation forward tunnel, got %d", runtimeCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationTunnelCreateRejectsOccupiedPort(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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, is_remote, remote_url, remote_token, remote_config)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "federation-port-check-node", "federation-port-check-secret", "10.91.80.70", "10.91.80.70", "", "24100-24120", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "federation-port-check-node")
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "federation-port-check-share",
|
||||
NodeID: nodeID,
|
||||
Token: "federation-port-check-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 24100,
|
||||
PortRangeEnd: 24120,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
|
||||
create := func() response.R {
|
||||
body, err := json.Marshal(federationTunnelRequest{Protocol: "tcp", RemotePort: 24101, Target: "1.1.1.1:443"})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/tunnel/create", bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", "Bearer federation-port-check-token")
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
h.federationTunnelCreate(res, req)
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
var payload response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
return payload
|
||||
}
|
||||
|
||||
first := create()
|
||||
if first.Code != 0 {
|
||||
t.Fatalf("expected first create success, got %d (%s)", first.Code, first.Msg)
|
||||
}
|
||||
|
||||
second := create()
|
||||
if second.Code != 403 {
|
||||
t.Fatalf("expected second create to be rejected with 403, got %d (%s)", second.Code, second.Msg)
|
||||
}
|
||||
if second.Msg != "Port already in use" {
|
||||
t.Fatalf("expected occupied port message, got %q", second.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteTunnelReleasesFederationForwardRuntimeByPort(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "delete-forward-share",
|
||||
NodeID: 1,
|
||||
Token: "delete-forward-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 25000,
|
||||
PortRangeEnd: 25020,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("delete-forward-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
tunnelName := fmt.Sprintf("Share-%d-Port-%d", share.ID, 25001)
|
||||
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, tunnelName, 1.0, 1, "tcp", 1, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, share.NodeID, "del-r1", "del-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 25001, "", 1, 1, now, now).Error; err != nil {
|
||||
t.Fatalf("insert runtime: %v", err)
|
||||
}
|
||||
|
||||
if err := h.deleteTunnelByID(1); err != nil {
|
||||
t.Fatalf("delete tunnel: %v", err)
|
||||
}
|
||||
|
||||
activeCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND port = ? AND status = 1`, share.ID, 25001)
|
||||
if activeCount != 0 {
|
||||
t.Fatalf("expected runtime released after tunnel delete, active rows=%d", activeCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBindPeerShareForwardRuntimeServicesOnlyBindsForwardRole(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "bind-forward-role-share",
|
||||
NodeID: 1,
|
||||
Token: "bind-forward-role-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 26000,
|
||||
PortRangeEnd: 26020,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("bind-forward-role-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?),
|
||||
(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`,
|
||||
1, share.ID, share.NodeID, "bind-r1", "bind-rk1", "", "forward", "", "", "tcp", "fifo", 26001, "", 0, 1, now, now,
|
||||
2, share.ID, share.NodeID, "bind-r2", "bind-rk2", "", "middle", "", "", "tcp", "round", 26002, "", 0, 1, now, now,
|
||||
).Error; err != nil {
|
||||
t.Fatalf("insert runtimes: %v", err)
|
||||
}
|
||||
|
||||
h.bindPeerShareForwardRuntimeServices(share, map[string]interface{}{
|
||||
"services": []interface{}{
|
||||
map[string]interface{}{"name": "77_2_10_tcp", "addr": "[::]:26001"},
|
||||
map[string]interface{}{"name": "88_2_10_tcp", "addr": "[::]:26002"},
|
||||
},
|
||||
})
|
||||
|
||||
forwardServiceName := ""
|
||||
middleServiceName := ""
|
||||
if err := r.DB().Raw(`SELECT service_name FROM peer_share_runtime WHERE id = 1`).Scan(&forwardServiceName).Error; err != nil {
|
||||
t.Fatalf("load forward runtime service name: %v", err)
|
||||
}
|
||||
if err := r.DB().Raw(`SELECT service_name FROM peer_share_runtime WHERE id = 2`).Scan(&middleServiceName).Error; err != nil {
|
||||
t.Fatalf("load middle runtime service name: %v", err)
|
||||
}
|
||||
|
||||
if forwardServiceName != "77_2_10" {
|
||||
t.Fatalf("expected forward runtime service name bound, got %q", forwardServiceName)
|
||||
}
|
||||
if middleServiceName != "" {
|
||||
t.Fatalf("expected non-forward runtime unchanged, got %q", middleServiceName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBindPeerShareForwardRuntimeServicesAcceptsTopLevelServiceArray(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "bind-array-share",
|
||||
NodeID: 1,
|
||||
Token: "bind-array-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 26100,
|
||||
PortRangeEnd: 26120,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("bind-array-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`,
|
||||
1, share.ID, share.NodeID, "bind-array-r1", "bind-array-rk1", "", "forward", "", "", "tcp", "fifo", 26101, "", 0, 1, now, now,
|
||||
).Error; err != nil {
|
||||
t.Fatalf("insert runtime: %v", err)
|
||||
}
|
||||
|
||||
h.bindPeerShareForwardRuntimeServices(share, []interface{}{
|
||||
map[string]interface{}{"name": "99_2_10_tcp", "addr": "[::]:26101"},
|
||||
map[string]interface{}{"name": "99_2_10_udp", "addr": "[::]:26101"},
|
||||
})
|
||||
|
||||
forwardServiceName := ""
|
||||
if err := r.DB().Raw(`SELECT service_name FROM peer_share_runtime WHERE id = 1`).Scan(&forwardServiceName).Error; err != nil {
|
||||
t.Fatalf("load forward runtime service name: %v", err)
|
||||
}
|
||||
if forwardServiceName != "99_2_10" {
|
||||
t.Fatalf("expected forward runtime service name bound from top-level array, got %q", forwardServiceName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBindPeerShareForwardRuntimeServicesCreatesRuntimeWhenMissing(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-bind-create-runtime.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "bind-create-runtime-share",
|
||||
NodeID: 1,
|
||||
Token: "bind-create-runtime-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 26300,
|
||||
PortRangeEnd: 26320,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("bind-create-runtime-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
h.bindPeerShareForwardRuntimeServices(share, map[string]interface{}{
|
||||
"services": []interface{}{
|
||||
map[string]interface{}{"name": "55_2_10_tcp", "addr": "[::]:26301"},
|
||||
},
|
||||
})
|
||||
|
||||
var count int64
|
||||
if err := r.DB().Raw(`SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND role = ? AND status = 1`, share.ID, "forward").Scan(&count).Error; err != nil {
|
||||
t.Fatalf("query runtime count: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Fatalf("expected 1 active forward runtime row, got %d", count)
|
||||
}
|
||||
|
||||
var serviceName string
|
||||
var port int
|
||||
var applied int
|
||||
if err := r.DB().Raw(`SELECT service_name, port, applied FROM peer_share_runtime WHERE share_id = ? AND role = ? ORDER BY id DESC LIMIT 1`, share.ID, "forward").Row().Scan(&serviceName, &port, &applied); err != nil {
|
||||
t.Fatalf("query created runtime: %v", err)
|
||||
}
|
||||
if serviceName != "55_2_10" {
|
||||
t.Fatalf("expected service_name=55_2_10, got %q", serviceName)
|
||||
}
|
||||
if port != 26301 {
|
||||
t.Fatalf("expected port=26301, got %d", port)
|
||||
}
|
||||
if applied != 1 {
|
||||
t.Fatalf("expected applied=1, got %d", applied)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleasePeerShareForwardRuntimeServicesMarksRuntimeReleased(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-release-runtime.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "release-runtime-share",
|
||||
NodeID: 1,
|
||||
Token: "release-runtime-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 26400,
|
||||
PortRangeEnd: 26420,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("release-runtime-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, share.NodeID, "release-r1", "release-rk1", "", "forward", "", "77_2_10", "tcp", "fifo", 26401, "", 1, 1, now, now).Error; err != nil {
|
||||
t.Fatalf("insert runtime: %v", err)
|
||||
}
|
||||
|
||||
h.releasePeerShareForwardRuntimeServices(share, map[string]interface{}{
|
||||
"services": []interface{}{"77_2_10_tcp"},
|
||||
})
|
||||
|
||||
var status int
|
||||
var applied int
|
||||
var serviceName string
|
||||
if err := r.DB().Raw(`SELECT status, applied, service_name FROM peer_share_runtime WHERE share_id = ? AND role = ? ORDER BY id DESC LIMIT 1`, share.ID, "forward").Row().Scan(&status, &applied, &serviceName); err != nil {
|
||||
t.Fatalf("query released runtime: %v", err)
|
||||
}
|
||||
if status != 0 {
|
||||
t.Fatalf("expected status=0 after release, got %d", status)
|
||||
}
|
||||
if applied != 0 {
|
||||
t.Fatalf("expected applied=0 after release, got %d", applied)
|
||||
}
|
||||
if serviceName != "" {
|
||||
t.Fatalf("expected service_name cleared after release, got %q", serviceName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateFederationCommandPortsAcceptsTopLevelServiceArray(t *testing.T) {
|
||||
share := &repo.PeerShare{
|
||||
PortRangeStart: 26200,
|
||||
PortRangeEnd: 26210,
|
||||
}
|
||||
err := validateFederationCommandPorts(share, []interface{}{
|
||||
map[string]interface{}{"name": "11_2_10_tcp", "addr": "[::]:26201"},
|
||||
map[string]interface{}{"name": "11_2_10_udp", "addr": "[::]:26201"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("expected top-level service array to pass port validation, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRemoteUsageList(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
@@ -502,6 +941,110 @@ func TestFederationRemoteUsageList(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRemoteUsageListIncludesForwardPorts(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-forward-usage.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "test-jwt-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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, is_remote, remote_url, remote_token, remote_config)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-usage-remote-node", "forward-usage-secret", "10.60.70.80", "10.60.70.80", "", "33000-33010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "", "", `{"shareId":99,"maxBandwidth":0,"currentFlow":0,"portRangeStart":33000,"portRangeEnd":33010}`).Error; err != nil {
|
||||
t.Fatalf("insert remote node: %v", err)
|
||||
}
|
||||
|
||||
var nodeID int64
|
||||
if err := r.DB().Raw(`SELECT id FROM node WHERE name = ? ORDER BY id DESC LIMIT 1`, "forward-usage-remote-node").Row().Scan(&nodeID); err != nil {
|
||||
t.Fatalf("query node id: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-usage-tunnel", 1, "tls", 1, now, now, 1, "", 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
var tunnelID int64
|
||||
if err := r.DB().Raw(`SELECT id FROM tunnel WHERE name = ? ORDER BY id DESC LIMIT 1`, "forward-usage-tunnel").Row().Scan(&tunnelID); err != nil {
|
||||
t.Fatalf("query tunnel id: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "tester", "forward-usage-item", tunnelID, "1.1.1.1:443", "fifo", 0, 0, now, now, 1, 0).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
var forwardID int64
|
||||
if err := r.DB().Raw(`SELECT id FROM forward WHERE name = ? ORDER BY id DESC LIMIT 1`, "forward-usage-item").Row().Scan(&forwardID); err != nil {
|
||||
t.Fatalf("query forward id: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, 33001).Error; err != nil {
|
||||
t.Fatalf("insert forward_port: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/federation/share/remote-usage/list", nil)
|
||||
res := httptest.NewRecorder()
|
||||
h.federationRemoteUsageList(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
|
||||
var payload response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected response code 0, got %d (%s)", payload.Code, payload.Msg)
|
||||
}
|
||||
|
||||
rows, ok := payload.Data.([]interface{})
|
||||
if !ok || len(rows) == 0 {
|
||||
t.Fatalf("expected non-empty usage list, got %T", payload.Data)
|
||||
}
|
||||
|
||||
first, ok := rows[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected usage row map, got %T", rows[0])
|
||||
}
|
||||
|
||||
usedPortsRaw, ok := first["usedPorts"].([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected usedPorts array, got %T", first["usedPorts"])
|
||||
}
|
||||
if len(usedPortsRaw) != 1 || int(usedPortsRaw[0].(float64)) != 33001 {
|
||||
t.Fatalf("expected usedPorts [33001], got %v", usedPortsRaw)
|
||||
}
|
||||
|
||||
bindingsRaw, ok := first["bindings"].([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected bindings array, got %T", first["bindings"])
|
||||
}
|
||||
if len(bindingsRaw) != 1 {
|
||||
t.Fatalf("expected 1 binding row from forward usage, got %d", len(bindingsRaw))
|
||||
}
|
||||
|
||||
binding, ok := bindingsRaw[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected binding row object, got %T", bindingsRaw[0])
|
||||
}
|
||||
if int(binding["allocatedPort"].(float64)) != 33001 {
|
||||
t.Fatalf("expected allocatedPort=33001, got %v", binding["allocatedPort"])
|
||||
}
|
||||
if int(binding["chainType"].(float64)) != 1 {
|
||||
t.Fatalf("expected chainType=1 for forward usage row, got %v", binding["chainType"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthPeerAllowedIPs(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
|
||||
@@ -2,9 +2,13 @@ package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
const bytesPerGB int64 = 1024 * 1024 * 1024
|
||||
@@ -18,6 +22,7 @@ type userTunnelPolicy struct {
|
||||
OutFlow int64
|
||||
ExpTime int64
|
||||
Status int
|
||||
Num int
|
||||
}
|
||||
|
||||
type gostConfigSnapshot struct {
|
||||
@@ -30,7 +35,7 @@ type namedConfigItem struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
func (h *Handler) processFlowItem(item flowItem) {
|
||||
func (h *Handler) processFlowItem(nodeID int64, item flowItem) {
|
||||
serviceName := strings.TrimSpace(item.N)
|
||||
if serviceName == "" || serviceName == "web_api" {
|
||||
return
|
||||
@@ -40,6 +45,10 @@ func (h *Handler) processFlowItem(item flowItem) {
|
||||
if ok {
|
||||
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
|
||||
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
|
||||
if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil {
|
||||
h.enforceUserQuotaIfNeeded(userID, quota)
|
||||
}
|
||||
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
|
||||
|
||||
if userTunnelID > 0 {
|
||||
h.enforceFlowPolicies(userID, userTunnelID)
|
||||
@@ -87,6 +96,45 @@ func parsePeerShareRuntimeServiceID(serviceName string) (int64, bool) {
|
||||
return runtimeID, true
|
||||
}
|
||||
|
||||
func parsePeerShareInfoFromFederationTunnelName(tunnelName string) (int64, int, bool) {
|
||||
tunnelName = strings.TrimSpace(tunnelName)
|
||||
if !strings.HasPrefix(tunnelName, "Share-") {
|
||||
return 0, 0, false
|
||||
}
|
||||
raw := strings.TrimPrefix(tunnelName, "Share-")
|
||||
idx := strings.Index(raw, "-Port-")
|
||||
if idx <= 0 {
|
||||
return 0, 0, false
|
||||
}
|
||||
shareID, err := strconv.ParseInt(raw[:idx], 10, 64)
|
||||
if err != nil || shareID <= 0 {
|
||||
return 0, 0, false
|
||||
}
|
||||
portValue := strings.TrimSpace(raw[idx+len("-Port-"):])
|
||||
port, err := strconv.Atoi(portValue)
|
||||
if err != nil || port <= 0 {
|
||||
return 0, 0, false
|
||||
}
|
||||
return shareID, port, true
|
||||
}
|
||||
|
||||
func parsePeerShareIDFromFederationTunnelName(tunnelName string) (int64, bool) {
|
||||
tunnelName = strings.TrimSpace(tunnelName)
|
||||
if !strings.HasPrefix(tunnelName, "Share-") {
|
||||
return 0, false
|
||||
}
|
||||
raw := strings.TrimPrefix(tunnelName, "Share-")
|
||||
idx := strings.Index(raw, "-Port-")
|
||||
if idx <= 0 {
|
||||
return 0, false
|
||||
}
|
||||
shareID, err := strconv.ParseInt(raw[:idx], 10, 64)
|
||||
if err != nil || shareID <= 0 {
|
||||
return 0, false
|
||||
}
|
||||
return shareID, true
|
||||
}
|
||||
|
||||
func (h *Handler) processPeerShareFlow(runtimeID int64, item flowItem) {
|
||||
if h == nil || h.repo == nil || runtimeID <= 0 {
|
||||
return
|
||||
@@ -113,6 +161,121 @@ func (h *Handler) processPeerShareFlow(runtimeID int64, item flowItem) {
|
||||
h.enforcePeerShareFlowLimit(share.ID)
|
||||
}
|
||||
|
||||
func (h *Handler) processPeerShareFlowFromForward(forwardID int64, nodeID int64, serviceName string, item flowItem) {
|
||||
if h == nil || h.repo == nil || forwardID <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
delta := item.D + item.U
|
||||
if delta <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil || forward == nil {
|
||||
// Forward not found in local database - might be a federation port-forward
|
||||
// Try to find by service name in peer_share_runtime
|
||||
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
|
||||
return
|
||||
}
|
||||
tunnelName, err := h.repo.GetTunnelName(forward.TunnelID)
|
||||
if err != nil {
|
||||
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
|
||||
return
|
||||
}
|
||||
shareID, ok := parsePeerShareIDFromFederationTunnelName(tunnelName)
|
||||
if !ok {
|
||||
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.AddPeerShareCurrentFlow(shareID, delta); err != nil {
|
||||
h.processPeerShareFlowByServiceName(nodeID, serviceName, item)
|
||||
return
|
||||
}
|
||||
|
||||
share, err := h.repo.GetPeerShare(shareID)
|
||||
if err != nil || share == nil {
|
||||
return
|
||||
}
|
||||
if !isPeerShareFlowExceeded(share) {
|
||||
return
|
||||
}
|
||||
h.enforcePeerShareFlowLimit(share.ID)
|
||||
}
|
||||
|
||||
func normalizeForwardRuntimeServiceName(serviceName string) string {
|
||||
name := strings.TrimSpace(serviceName)
|
||||
if strings.HasSuffix(name, "_tcp") {
|
||||
return strings.TrimSuffix(name, "_tcp")
|
||||
}
|
||||
if strings.HasSuffix(name, "_udp") {
|
||||
return strings.TrimSuffix(name, "_udp")
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
func (h *Handler) processPeerShareFlowByServiceName(nodeID int64, serviceName string, item flowItem) {
|
||||
if h == nil || h.repo == nil || strings.TrimSpace(serviceName) == "" {
|
||||
return
|
||||
}
|
||||
|
||||
delta := item.D + item.U
|
||||
if delta <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
normalized := normalizeForwardRuntimeServiceName(serviceName)
|
||||
var runtimes []model.PeerShareRuntime
|
||||
var err error
|
||||
|
||||
// Try node-scoped query first if nodeID is valid
|
||||
if nodeID > 0 {
|
||||
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID, normalized)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(runtimes) == 0 && normalized != serviceName {
|
||||
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByNodeAndServiceName(nodeID, serviceName)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to global query if node-scoped query returned nothing or nodeID is invalid
|
||||
if len(runtimes) == 0 {
|
||||
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByServiceName(normalized)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(runtimes) == 0 && normalized != serviceName {
|
||||
runtimes, err = h.repo.ListActiveForwardPeerShareRuntimesByServiceName(serviceName)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if len(runtimes) != 1 {
|
||||
if len(runtimes) > 1 {
|
||||
log.Printf("WARN: ambiguous peer share runtime match for service=%s nodeID=%d count=%d", serviceName, nodeID, len(runtimes))
|
||||
}
|
||||
return
|
||||
}
|
||||
runtime := runtimes[0]
|
||||
|
||||
_ = h.repo.AddPeerShareCurrentFlow(runtime.ShareID, delta)
|
||||
|
||||
matchedShare, err := h.repo.GetPeerShare(runtime.ShareID)
|
||||
if err != nil || matchedShare == nil {
|
||||
return
|
||||
}
|
||||
if isPeerShareFlowExceeded(matchedShare) {
|
||||
h.enforcePeerShareFlowLimit(matchedShare.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) enforcePeerShareFlowLimit(shareID int64) {
|
||||
if h == nil || h.repo == nil || shareID <= 0 {
|
||||
return
|
||||
@@ -169,6 +332,90 @@ func (h *Handler) enforceFlowPolicies(userID int64, userTunnelID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, now int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("invalid flow policy context")
|
||||
}
|
||||
if userID <= 0 || tunnelID <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
user, err := h.repo.GetUserByID(userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if user == nil {
|
||||
return errors.New("用户不存在")
|
||||
}
|
||||
|
||||
if user.Status != 1 {
|
||||
return errors.New("账号已禁用")
|
||||
}
|
||||
if user.ExpTime > 0 && user.ExpTime <= now {
|
||||
return errors.New("账号已过期")
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return errors.New("流量已超额,禁止开启转发")
|
||||
}
|
||||
if err := h.ensureUserForwardAllowedByQuota(userID, now); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if user.Num > 0 {
|
||||
currentForwardCount, err := h.repo.CountActiveForwardsByUser(userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if currentForwardCount >= int64(user.Num) {
|
||||
return errors.New("转发数量已达上限")
|
||||
}
|
||||
}
|
||||
|
||||
userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(userID, tunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if userTunnelID <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
policy, err := h.getUserTunnelPolicy(userTunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if policy == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if policy.Status != 1 {
|
||||
return errors.New("该隧道已禁用")
|
||||
}
|
||||
if policy.ExpTime > 0 && policy.ExpTime <= now {
|
||||
return errors.New("该隧道已过期")
|
||||
}
|
||||
|
||||
utFlowLimit := policy.Flow * bytesPerGB
|
||||
utCurrent := policy.InFlow + policy.OutFlow
|
||||
if utCurrent >= utFlowLimit {
|
||||
return errors.New("该隧道流量已超额,禁止开启转发")
|
||||
}
|
||||
|
||||
if policy.Num > 0 {
|
||||
currentTunnelForwardCount, err := h.repo.CountActiveForwardsByUserTunnel(userID, tunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if currentTunnelForwardCount >= int64(policy.Num) {
|
||||
return errors.New("该隧道转发数量已达上限")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
|
||||
user, err := h.repo.GetUserByID(userID)
|
||||
if err != nil || user == nil {
|
||||
@@ -216,7 +463,7 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er
|
||||
return &userTunnelPolicy{
|
||||
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
|
||||
Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
ExpTime: ut.ExpTime, Status: ut.Status,
|
||||
ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -271,15 +518,46 @@ func (h *Handler) cleanNodeConfigs(nodeID int64, rawConfig string) {
|
||||
}
|
||||
|
||||
func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem) {
|
||||
runtimeServiceNames, err := h.repo.ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
minUpdatedTime := time.Now().Add(-10 * time.Minute).UnixMilli()
|
||||
hasUnboundForwardPeerRuntime, err := h.repo.HasRecentUnboundForwardPeerShareRuntimeOnNode(nodeID, minUpdatedTime)
|
||||
if err != nil {
|
||||
hasUnboundForwardPeerRuntime = false
|
||||
}
|
||||
runtimeServiceSet := make(map[string]struct{}, len(runtimeServiceNames))
|
||||
for _, serviceName := range runtimeServiceNames {
|
||||
serviceName = strings.TrimSpace(serviceName)
|
||||
if serviceName == "" {
|
||||
continue
|
||||
}
|
||||
runtimeServiceSet[serviceName] = struct{}{}
|
||||
}
|
||||
|
||||
for _, item := range services {
|
||||
name := strings.TrimSpace(item.Name)
|
||||
if name == "" || name == "web_api" {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(name, "fed_svc_") {
|
||||
continue
|
||||
}
|
||||
normalizedName := normalizeForwardRuntimeServiceName(name)
|
||||
if _, ok := runtimeServiceSet[normalizedName]; ok {
|
||||
continue
|
||||
}
|
||||
if _, ok := runtimeServiceSet[name]; ok {
|
||||
continue
|
||||
}
|
||||
|
||||
parts := strings.Split(name, "_")
|
||||
if len(parts) >= 3 {
|
||||
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
|
||||
if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime {
|
||||
continue
|
||||
}
|
||||
if err == nil && forwardID > 0 && !h.forwardExists(forwardID) {
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name, parts[0] + "_" + parts[1] + "_" + parts[2], parts[0] + "_" + parts[1] + "_" + parts[2] + "_tcp", parts[0] + "_" + parts[1] + "_" + parts[2] + "_udp"}}, false, true)
|
||||
continue
|
||||
@@ -299,6 +577,9 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem
|
||||
continue
|
||||
}
|
||||
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
|
||||
if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime {
|
||||
continue
|
||||
}
|
||||
if err != nil || forwardID <= 0 || h.forwardExists(forwardID) {
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -43,7 +44,7 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) {
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.processFlowItem(flowItem{N: "fed_svc_17", U: 1200, D: 900})
|
||||
h.processFlowItem(1, flowItem{N: "fed_svc_17", U: 1200, D: 900})
|
||||
|
||||
updatedShare, err := r.GetPeerShare(share.ID)
|
||||
if err != nil || updatedShare == nil {
|
||||
@@ -61,3 +62,345 @@ func TestProcessFlowItemTracksPeerShareFlowAndEnforcesLimit(t *testing.T) {
|
||||
t.Fatalf("expected runtime status=0 after limit enforcement, got %d", runtime.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessFlowItemTracksPeerShareFlowForFederationPortForward(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-forward.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "forward-share",
|
||||
NodeID: 1,
|
||||
Token: "forward-share-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 30000,
|
||||
PortRangeEnd: 30010,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("forward-share-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'u2', 'x', 1, ?, 99999, 0, 0, 1, 1, ?, ?, 1)
|
||||
`, now+24*60*60*1000, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
tunnelName := "Share-" + strconv.FormatInt(share.ID, 10) + "-Port-30001"
|
||||
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, ?, 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, tunnelName, 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, 1, 99999, 0, 0, 1, ?, 1)
|
||||
`, now+24*60*60*1000).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:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.processFlowItem(1, flowItem{N: "20_2_10", U: 120, D: 80})
|
||||
|
||||
updatedShare, err := r.GetPeerShare(share.ID)
|
||||
if err != nil || updatedShare == nil {
|
||||
t.Fatalf("reload share: %v", err)
|
||||
}
|
||||
if updatedShare.CurrentFlow != 200 {
|
||||
t.Fatalf("expected current_flow=200, got %d", updatedShare.CurrentFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessFlowItemTracksPeerShareFlowByForwardServiceName(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-forward-service.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "forward-service-share",
|
||||
NodeID: 1,
|
||||
Token: "forward-service-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 31000,
|
||||
PortRangeEnd: 31010,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("forward-service-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, share.NodeID, "svc-r1", "svc-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31001, "", 1, 1, now, now).Error; err != nil {
|
||||
t.Fatalf("insert peer_share_runtime: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.processFlowItem(1, flowItem{N: "20_2_10_tcp", U: 120, D: 80})
|
||||
|
||||
updatedShare, err := r.GetPeerShare(share.ID)
|
||||
if err != nil || updatedShare == nil {
|
||||
t.Fatalf("reload share: %v", err)
|
||||
}
|
||||
if updatedShare.CurrentFlow != 200 {
|
||||
t.Fatalf("expected current_flow=200, got %d", updatedShare.CurrentFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessFlowItemFallsBackToServiceNameWhenForwardIDCollidesAcrossPanels(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-forward-collision.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "collision-share",
|
||||
NodeID: 1,
|
||||
Token: "collision-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 31400,
|
||||
PortRangeEnd: 31410,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("collision-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, share.NodeID, "collision-r1", "collision-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31401, "", 1, 1, now, now).Error; err != nil {
|
||||
t.Fatalf("insert peer_share_runtime: %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(2, 'local-tunnel-with-colliding-forward-id', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert local 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, 1, 'local-user', 'local-f20', 2, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert local forward: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.processFlowItem(1, flowItem{N: "20_2_10_tcp", U: 120, D: 80})
|
||||
|
||||
updatedShare, err := r.GetPeerShare(share.ID)
|
||||
if err != nil || updatedShare == nil {
|
||||
t.Fatalf("reload share: %v", err)
|
||||
}
|
||||
if updatedShare.CurrentFlow != 200 {
|
||||
t.Fatalf("expected current_flow=200, got %d", updatedShare.CurrentFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessFlowItemSkipsPeerShareFlowWhenServiceNameIsAmbiguous(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-forward-ambiguous.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "ambiguous-share-a",
|
||||
NodeID: 1,
|
||||
Token: "ambiguous-token-a",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 31100,
|
||||
PortRangeEnd: 31110,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share A: %v", err)
|
||||
}
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "ambiguous-share-b",
|
||||
NodeID: 1,
|
||||
Token: "ambiguous-token-b",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 31200,
|
||||
PortRangeEnd: 31210,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share B: %v", err)
|
||||
}
|
||||
shareA, _ := r.GetPeerShareByToken("ambiguous-token-a")
|
||||
shareB, _ := r.GetPeerShareByToken("ambiguous-token-b")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?),
|
||||
(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`,
|
||||
shareA.ID, 1, "amb-r1", "amb-rk1", "", "forward", "", "99_2_10", "tcp", "fifo", 31101, "", 1, 1, now, now,
|
||||
shareB.ID, 1, "amb-r2", "amb-rk2", "", "forward", "", "99_2_10", "tcp", "fifo", 31201, "", 1, 1, now, now,
|
||||
).Error; err != nil {
|
||||
t.Fatalf("insert ambiguous runtimes: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.processFlowItem(1, flowItem{N: "99_2_10_tcp", U: 120, D: 80})
|
||||
|
||||
updatedA, _ := r.GetPeerShare(shareA.ID)
|
||||
updatedB, _ := r.GetPeerShare(shareB.ID)
|
||||
if updatedA.CurrentFlow != 0 || updatedB.CurrentFlow != 0 {
|
||||
t.Fatalf("expected ambiguous service flow to be skipped, got shareA=%d shareB=%d", updatedA.CurrentFlow, updatedB.CurrentFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanOrphanedServicesSkipsActiveSharedForwardRuntimeServices(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-cleanup-runtime.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "cleanup-runtime-share",
|
||||
NodeID: 1,
|
||||
Token: "cleanup-runtime-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 31300,
|
||||
PortRangeEnd: 31310,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("cleanup-runtime-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, share.NodeID, "cleanup-r1", "cleanup-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31301, "", 1, 1, now, now).Error; err != nil {
|
||||
t.Fatalf("insert peer_share_runtime: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
t.Fatalf("cleanOrphanedServices should skip active shared runtime service; got panic: %v", rec)
|
||||
}
|
||||
}()
|
||||
|
||||
h.cleanOrphanedServices(share.NodeID, []namedConfigItem{{Name: "20_2_10_tcp"}})
|
||||
}
|
||||
|
||||
func TestCleanOrphanedServicesSkipsFederationServicePrefix(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-cleanup-fed-svc.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
h := &Handler{repo: r}
|
||||
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
t.Fatalf("cleanOrphanedServices should skip fed_svc_ service names; got panic: %v", rec)
|
||||
}
|
||||
}()
|
||||
|
||||
h.cleanOrphanedServices(1, []namedConfigItem{{Name: "fed_svc_999_tcp"}})
|
||||
}
|
||||
|
||||
func TestCleanOrphanedServicesSkipsForwardPatternWhenNodeHasActivePeerShareForwardRuntime(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel-cleanup-forward-runtime-empty-service.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "cleanup-forward-runtime-empty-service",
|
||||
NodeID: 1,
|
||||
Token: "cleanup-forward-runtime-empty-service-token",
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 0,
|
||||
PortRangeStart: 31420,
|
||||
PortRangeEnd: 31430,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create peer share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken("cleanup-forward-runtime-empty-service-token")
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load peer share: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, share.ID, share.NodeID, "cleanup-forward-empty-r1", "cleanup-forward-empty-rk1", "", "forward", "", "", "tcp", "fifo", 31421, "", 0, 1, now, now).Error; err != nil {
|
||||
t.Fatalf("insert peer_share_runtime with empty service name: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
t.Fatalf("cleanOrphanedServices should skip forward-pattern services when active peer-share forward runtime exists; got panic: %v", rec)
|
||||
}
|
||||
}()
|
||||
|
||||
h.cleanOrphanedServices(share.NodeID, []namedConfigItem{{Name: "20_2_10_tcp"}})
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -15,17 +16,21 @@ import (
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/health"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
)
|
||||
|
||||
type Handler struct {
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
|
||||
captchaMu sync.Mutex
|
||||
captchaTokens map[string]int64
|
||||
@@ -34,8 +39,15 @@ type Handler struct {
|
||||
jobsCancel context.CancelFunc
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
|
||||
upgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
|
||||
qualityProber *tunnelQualityProber
|
||||
}
|
||||
|
||||
const monitorTunnelQualityEnabledConfigKey = "monitor_tunnel_quality_enabled"
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
@@ -69,13 +81,43 @@ type flowItem struct {
|
||||
D int64 `json:"d"`
|
||||
}
|
||||
|
||||
const (
|
||||
pngDataURLPrefix = "data:image/png;base64,"
|
||||
maxBrandAssetDataURLBytes = 1024 * 1024
|
||||
)
|
||||
|
||||
func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
return &Handler{
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
captchaTokens: make(map[string]int64),
|
||||
h := &Handler{
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
}
|
||||
h.healthCheck = health.NewChecker(repo, h.wsServer)
|
||||
h.qualityProber = newTunnelQualityProber(h)
|
||||
h.wsServer.SetNodeOnlineHook(h.onNodeOnline)
|
||||
h.wsServer.SetNodeMetricHook(func(nodeID int64, info ws.SystemInfo) {
|
||||
metricInfo := metrics.SystemInfo{
|
||||
Uptime: info.Uptime,
|
||||
BytesReceived: info.BytesReceived,
|
||||
BytesTransmitted: info.BytesTransmitted,
|
||||
CPUUsage: info.CPUUsage,
|
||||
MemoryUsage: info.MemoryUsage,
|
||||
DiskUsage: info.DiskUsage,
|
||||
Load1: info.Load1,
|
||||
Load5: info.Load5,
|
||||
Load15: info.Load15,
|
||||
TCPConns: info.TCPConns,
|
||||
UDPConns: info.UDPConns,
|
||||
NetInSpeed: info.NetInSpeed,
|
||||
NetOutSpeed: info.NetOutSpeed,
|
||||
}
|
||||
h.metrics.RecordNodeMetric(nodeID, metricInfo)
|
||||
})
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *Handler) WebSocketHandler() http.Handler {
|
||||
@@ -89,6 +131,8 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/user/update", h.userUpdate)
|
||||
mux.HandleFunc("/api/v1/user/delete", h.userDelete)
|
||||
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/config/get", h.getConfigByName)
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||
@@ -109,6 +153,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/node/delete", h.nodeDelete)
|
||||
mux.HandleFunc("/api/v1/node/install", h.nodeInstall)
|
||||
mux.HandleFunc("/api/v1/node/update-order", h.nodeUpdateOrder)
|
||||
mux.HandleFunc("/api/v1/node/dismiss-expiry-reminder", h.nodeDismissExpiryReminder)
|
||||
mux.HandleFunc("/api/v1/node/batch-delete", h.nodeBatchDelete)
|
||||
mux.HandleFunc("/api/v1/node/check-status", h.nodeCheckStatus)
|
||||
mux.HandleFunc("/api/v1/node/upgrade", h.nodeUpgrade)
|
||||
@@ -120,7 +165,12 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
mux.HandleFunc("/api/v1/tunnel/update", h.tunnelUpdate)
|
||||
mux.HandleFunc("/api/v1/tunnel/delete", h.tunnelDelete)
|
||||
mux.HandleFunc("/api/v1/tunnel/delete-preview", h.tunnelDeletePreview)
|
||||
mux.HandleFunc("/api/v1/tunnel/delete-with-forwards", h.tunnelDeleteWithForwards)
|
||||
mux.HandleFunc("/api/v1/tunnel/batch-delete-preview", h.tunnelBatchDeletePreview)
|
||||
mux.HandleFunc("/api/v1/tunnel/batch-delete-with-forwards", h.tunnelBatchDeleteWithForwards)
|
||||
mux.HandleFunc("/api/v1/tunnel/diagnose", h.tunnelDiagnose)
|
||||
mux.HandleFunc("/api/v1/tunnel/diagnose/stream", h.tunnelDiagnoseStream)
|
||||
mux.HandleFunc("/api/v1/tunnel/update-order", h.tunnelUpdateOrder)
|
||||
mux.HandleFunc("/api/v1/tunnel/batch-delete", h.tunnelBatchDelete)
|
||||
mux.HandleFunc("/api/v1/tunnel/batch-redeploy", h.tunnelBatchRedeploy)
|
||||
@@ -136,6 +186,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/forward/pause", h.forwardPause)
|
||||
mux.HandleFunc("/api/v1/forward/resume", h.forwardResume)
|
||||
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)
|
||||
mux.HandleFunc("/api/v1/forward/batch-delete", h.forwardBatchDelete)
|
||||
mux.HandleFunc("/api/v1/forward/batch-pause", h.forwardBatchPause)
|
||||
@@ -146,7 +197,6 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/speed-limit/create", h.speedLimitCreate)
|
||||
mux.HandleFunc("/api/v1/speed-limit/update", h.speedLimitUpdate)
|
||||
mux.HandleFunc("/api/v1/speed-limit/delete", h.speedLimitDelete)
|
||||
mux.HandleFunc("/api/v1/speed-limit/tunnels", h.tunnelList)
|
||||
mux.HandleFunc("/api/v1/tunnel/user/tunnel", h.userTunnelVisibleList)
|
||||
mux.HandleFunc("/api/v1/tunnel/user/list", h.userTunnelList)
|
||||
mux.HandleFunc("/api/v1/group/tunnel/list", h.tunnelGroupList)
|
||||
@@ -180,6 +230,24 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/announcement/get", h.getAnnouncement)
|
||||
mux.HandleFunc("/api/v1/announcement/update", h.updateAnnouncement)
|
||||
|
||||
mux.HandleFunc("/api/v1/monitor/access", h.monitorAccessHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/nodes/", h.monitorNodeMetricsHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/nodes", h.monitorNodeListHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/tunnels", h.monitorTunnelListHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/tunnels/quality", h.monitorTunnelQualityHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/tunnels/", h.monitorTunnelMetrics)
|
||||
mux.HandleFunc("/api/v1/monitor/services", h.monitorServiceListHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/services/create", h.monitorServiceCreate)
|
||||
mux.HandleFunc("/api/v1/monitor/services/update", h.monitorServiceUpdate)
|
||||
mux.HandleFunc("/api/v1/monitor/services/delete", h.monitorServiceDelete)
|
||||
mux.HandleFunc("/api/v1/monitor/services/run", h.monitorServiceRun)
|
||||
mux.HandleFunc("/api/v1/monitor/services/latest-results", h.monitorServiceLatestResultsHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/services/limits", h.monitorServiceLimitsHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/services/", h.monitorServiceResultsHandler)
|
||||
mux.HandleFunc("/api/v1/monitor/permission/list", h.monitorPermissionList)
|
||||
mux.HandleFunc("/api/v1/monitor/permission/assign", h.monitorPermissionAssign)
|
||||
mux.HandleFunc("/api/v1/monitor/permission/remove", h.monitorPermissionRemove)
|
||||
|
||||
mux.HandleFunc("/flow/test", h.flowTest)
|
||||
mux.HandleFunc("/flow/config", h.flowConfig)
|
||||
mux.HandleFunc("/flow/upload", h.flowUpload)
|
||||
@@ -212,7 +280,7 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if captchaEnabled {
|
||||
if captchaEnabled && !h.apiClientCaptchaBypassEnabled(r) {
|
||||
captchaID := strings.TrimSpace(req.CaptchaID)
|
||||
if captchaID == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
|
||||
@@ -557,7 +625,7 @@ func (h *Handler) userTunnelList(w http.ResponseWriter, r *http.Request) {
|
||||
"userId": t.UserID,
|
||||
"tunnelId": t.TunnelID,
|
||||
"tunnelName": t.TunnelName,
|
||||
"status": 1,
|
||||
"status": t.Status,
|
||||
"flow": t.Flow,
|
||||
"num": t.Num,
|
||||
"expTime": t.ExpTime,
|
||||
@@ -697,7 +765,8 @@ func (h *Handler) flowConfig(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
secret := r.URL.Query().Get("secret")
|
||||
if ok, _ := h.repo.NodeExistsBySecret(secret); !ok {
|
||||
node, _ := h.repo.GetNodeBySecret(secret)
|
||||
if node == nil {
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
return
|
||||
@@ -707,8 +776,10 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
if err == nil && strings.TrimSpace(raw) != "" {
|
||||
var items []flowItem
|
||||
if json.Unmarshal([]byte(raw), &items) == nil {
|
||||
nowMs := time.Now().UnixMilli()
|
||||
h.recordTunnelMetricsFromFlowItems(node.ID, items, nowMs)
|
||||
for _, item := range items {
|
||||
h.processFlowItem(item)
|
||||
h.processFlowItem(node.ID, item)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -739,7 +810,14 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
if key == "" {
|
||||
continue
|
||||
}
|
||||
if err := h.repo.UpsertConfig(key, v, now); err != nil {
|
||||
|
||||
value, err := normalizeAndValidateConfigValue(key, v)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpsertConfig(key, value, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -759,16 +837,24 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(req.Name) == "" {
|
||||
name := strings.TrimSpace(req.Name)
|
||||
if name == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("配置名称不能为空"))
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(req.Value) == "" {
|
||||
|
||||
value, err := normalizeAndValidateConfigValue(name, req.Value)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if value == "" && name != "app_logo" && name != "app_favicon" {
|
||||
response.WriteJSON(w, response.ErrDefault("配置值不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpsertConfig(strings.TrimSpace(req.Name), req.Value, time.Now().UnixMilli()); err != nil {
|
||||
if err := h.repo.UpsertConfig(name, value, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -776,6 +862,58 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func normalizeAndValidateConfigValue(key, value string) (string, error) {
|
||||
switch strings.TrimSpace(key) {
|
||||
case "app_logo", "app_favicon":
|
||||
normalized := strings.TrimSpace(value)
|
||||
if normalized == "" {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
if !strings.HasPrefix(normalized, pngDataURLPrefix) {
|
||||
return "", fmt.Errorf("品牌图片必须通过上传生成 PNG 数据")
|
||||
}
|
||||
|
||||
if len(normalized) > maxBrandAssetDataURLBytes {
|
||||
return "", fmt.Errorf("品牌图片过大,请上传更小图片")
|
||||
}
|
||||
|
||||
payload := strings.TrimSpace(strings.TrimPrefix(normalized, pngDataURLPrefix))
|
||||
if payload == "" {
|
||||
return "", fmt.Errorf("品牌图片数据不能为空")
|
||||
}
|
||||
|
||||
if _, err := base64.StdEncoding.DecodeString(payload); err != nil {
|
||||
return "", fmt.Errorf("品牌图片数据格式无效")
|
||||
}
|
||||
|
||||
return pngDataURLPrefix + payload, nil
|
||||
case monitorTunnelQualityEnabledConfigKey:
|
||||
normalized := strings.TrimSpace(strings.ToLower(value))
|
||||
switch normalized {
|
||||
case "true", "false":
|
||||
return normalized, nil
|
||||
default:
|
||||
return "", fmt.Errorf("隧道质量检测开关配置值无效")
|
||||
}
|
||||
default:
|
||||
return value, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(monitorTunnelQualityEnabledConfigKey)
|
||||
if err != nil || cfg == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
return strings.TrimSpace(strings.ToLower(cfg.Value)) != "false"
|
||||
}
|
||||
|
||||
func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -981,10 +1119,41 @@ func (h *Handler) captchaEnabled() (bool, error) {
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if cfg == nil {
|
||||
if cfg == nil || !strings.EqualFold(strings.TrimSpace(cfg.Value), "true") {
|
||||
return false, nil
|
||||
}
|
||||
return strings.EqualFold(cfg.Value, "true"), nil
|
||||
|
||||
siteCfg, err := h.repo.GetConfigByName("cloudflare_site_key")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if siteCfg == nil || strings.TrimSpace(siteCfg.Value) == "" {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
secretCfg, err := h.repo.GetConfigByName("cloudflare_secret_key")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if secretCfg == nil || strings.TrimSpace(secretCfg.Value) == "" {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (h *Handler) apiClientCaptchaBypassEnabled(r *http.Request) bool {
|
||||
if r == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
client := strings.ToLower(strings.TrimSpace(r.Header.Get("X-FLVX-API-Client")))
|
||||
switch client {
|
||||
case "whmcs", "whmcs-module":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) markCaptchaToken(token string) {
|
||||
@@ -1092,6 +1261,21 @@ func nullableNullInt64(v sql.NullInt64) interface{} {
|
||||
return nil
|
||||
}
|
||||
|
||||
// flowCryptoCache caches AES crypto instances by secret to avoid per-request SHA256+GCM init.
|
||||
var flowCryptoCache sync.Map
|
||||
|
||||
func getOrCreateFlowCrypto(secret string) *security.AESCrypto {
|
||||
if v, ok := flowCryptoCache.Load(secret); ok {
|
||||
return v.(*security.AESCrypto)
|
||||
}
|
||||
c, err := security.NewAESCrypto(secret)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
flowCryptoCache.Store(secret, c)
|
||||
return c
|
||||
}
|
||||
|
||||
func readAndDecryptFlowBody(body io.ReadCloser, secret string) (string, error) {
|
||||
defer body.Close()
|
||||
raw, err := io.ReadAll(body)
|
||||
@@ -1112,8 +1296,8 @@ func readAndDecryptFlowBody(body io.ReadCloser, secret string) (string, error) {
|
||||
return text, nil
|
||||
}
|
||||
|
||||
crypto, err := security.NewAESCrypto(secret)
|
||||
if err != nil {
|
||||
crypto := getOrCreateFlowCrypto(secret)
|
||||
if crypto == nil {
|
||||
return text, nil
|
||||
}
|
||||
plain, err := crypto.Decrypt(wrap.Data)
|
||||
|
||||
@@ -18,11 +18,15 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
h.jobsCancel = cancel
|
||||
h.jobsStarted = true
|
||||
h.jobsWG.Add(2)
|
||||
h.jobsWG.Add(6)
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
go h.runHourlyStatsLoop(ctx)
|
||||
go h.runDailyMaintenanceLoop(ctx)
|
||||
go h.runNodeRenewalCycleLoop(ctx)
|
||||
go h.runMetricsIngestion(ctx)
|
||||
go h.runHealthChecks(ctx)
|
||||
go h.runTunnelQualityProber(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) StopBackgroundJobs() {
|
||||
@@ -46,6 +50,29 @@ func (h *Handler) StopBackgroundJobs() {
|
||||
h.jobsWG.Wait()
|
||||
}
|
||||
|
||||
func (h *Handler) runMetricsIngestion(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
if h.metrics != nil {
|
||||
h.metrics.Start(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runHealthChecks(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
if h.healthCheck != nil {
|
||||
h.healthCheck.Start(ctx)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runTunnelQualityProber(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
if h == nil || h.qualityProber == nil || !h.isTunnelQualityMonitoringEnabled() {
|
||||
return
|
||||
}
|
||||
|
||||
h.qualityProber.Start(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
@@ -135,6 +162,7 @@ func (h *Handler) runResetAndExpiryJob(now time.Time) {
|
||||
}
|
||||
|
||||
h.resetMonthlyFlow(now)
|
||||
h.resetUserQuotaWindows(now)
|
||||
h.disableExpiredUsers(now.UnixMilli())
|
||||
h.disableExpiredUserTunnels(now.UnixMilli())
|
||||
}
|
||||
@@ -176,3 +204,39 @@ func (h *Handler) disableExpiredUserTunnels(nowMs int64) {
|
||||
_ = h.repo.DisableUserTunnel(item.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runNodeRenewalCycleLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
for {
|
||||
wait := durationUntilNextNodeRenewalCycle(time.Now())
|
||||
timer := time.NewTimer(wait)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if !timer.Stop() {
|
||||
<-timer.C
|
||||
}
|
||||
return
|
||||
case <-timer.C:
|
||||
h.runNodeRenewalCycleJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func durationUntilNextNodeRenewalCycle(now time.Time) time.Duration {
|
||||
next := now.Truncate(6 * time.Hour).Add(6 * time.Hour)
|
||||
return next.Sub(now)
|
||||
}
|
||||
|
||||
func (h *Handler) runNodeRenewalCycleJob(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
advanced, err := h.repo.AdvanceNodeRenewalCycles(now.UnixMilli())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
_ = advanced
|
||||
}
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestRunNodeRenewalCycleJob_AdvancesOverdueAnchorTimes(t *testing.T) {
|
||||
dbPath := t.TempDir() + "/renewal-test.db"
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
})
|
||||
|
||||
now := time.Date(2026, 3, 8, 12, 0, 0, 0, time.UTC)
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
nodeID := int64(101)
|
||||
err = r.DB().Exec(`
|
||||
INSERT INTO node (id, name, secret, server_ip, port, http, tls, socks, created_time, status, renewal_cycle, expiry_time)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, nodeID, "no-cycle-node", "test-secret", "192.168.1.1", "1000-65535", 1, 1, 1, nowMs, 1, "", nil).Error
|
||||
if err != nil {
|
||||
t.Fatalf("insert test node: %v", err)
|
||||
}
|
||||
|
||||
quarterNodeID := int64(102)
|
||||
err = r.DB().Exec(`
|
||||
INSERT INTO node (id, name, secret, server_ip, port, http, tls, socks, created_time, status, renewal_cycle, expiry_time)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, quarterNodeID, "quarter-node", "test-secret", "192.168.1.1", "1000-65535", 1, 1, 1, nowMs, 1, "quarter", now.AddDate(0, -4, 0).UnixMilli()).Error
|
||||
if err != nil {
|
||||
t.Fatalf("insert test node: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.runNodeRenewalCycleJob(now)
|
||||
|
||||
var anchor sql.NullInt64
|
||||
err = r.DB().Raw(`SELECT expiry_time FROM node WHERE id = ?`, quarterNodeID).Row().Scan(&anchor)
|
||||
if err != nil {
|
||||
t.Fatalf("query expiry_time: %v", err)
|
||||
}
|
||||
|
||||
expectedAnchor := now.AddDate(0, 2, 0).UnixMilli()
|
||||
if !anchor.Valid || anchor.Int64 != expectedAnchor {
|
||||
t.Fatalf("expected anchor %d (2026-05-08), got %d", expectedAnchor, anchor.Int64)
|
||||
}
|
||||
}
|
||||
@@ -69,6 +69,13 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
t.Fatalf("insert expired user: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(3, 'non_expiring_user', 'x', 1, 0, 100, 1000, 2000, 15, 1, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert non-expiring 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', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
@@ -83,6 +90,13 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
t.Fatalf("insert expired user_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(11, 3, 1, NULL, 1, 1, 300, 400, 15, 0, 1)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert non-expiring 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, 'expired_user', 'f1', 1, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
@@ -90,6 +104,13 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
t.Fatalf("insert forward: %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(21, 3, 'non_expiring_user', 'f2', 1, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert non-expiring forward: %v", err)
|
||||
}
|
||||
|
||||
h.runResetAndExpiryJob(now)
|
||||
|
||||
userIn, userOut, userStatus := mustQueryInt64Int64Int(t, r, `SELECT in_flow, out_flow, status FROM user WHERE id = 2`)
|
||||
@@ -106,4 +127,56 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
if forwardStatus != 0 {
|
||||
t.Fatalf("expected forward status=0 after expiry handling, got %d", forwardStatus)
|
||||
}
|
||||
|
||||
nonExpUserStatus := mustQueryInt(t, r, `SELECT status FROM user WHERE id = 3`)
|
||||
if nonExpUserStatus != 1 {
|
||||
t.Fatalf("expected non-expiring user to remain enabled, got status=%d", nonExpUserStatus)
|
||||
}
|
||||
|
||||
nonExpTunnelStatus := mustQueryInt(t, r, `SELECT status FROM user_tunnel WHERE id = 11`)
|
||||
if nonExpTunnelStatus != 1 {
|
||||
t.Fatalf("expected non-expiring user_tunnel to remain enabled, got status=%d", nonExpTunnelStatus)
|
||||
}
|
||||
|
||||
nonExpForwardStatus := mustQueryInt(t, r, `SELECT status FROM forward WHERE id = 21`)
|
||||
if nonExpForwardStatus != 1 {
|
||||
t.Fatalf("expected non-expiring forward to remain enabled, got status=%d", nonExpForwardStatus)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResetAndExpiryJobResetsUserQuotaAndUnblocksUser(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "jobs-quota-reset.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "secret")
|
||||
now := time.Date(2026, 3, 12, 0, 0, 5, 0, time.UTC)
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'quota-reset-user', 'x', 1, 0, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
|
||||
`, 11*int64(1024*1024*1024), 11*int64(1024*1024*1024), nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user quota: %v", err)
|
||||
}
|
||||
|
||||
h.runResetAndExpiryJob(now)
|
||||
|
||||
dailyUsed := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 2`)
|
||||
if dailyUsed != 0 {
|
||||
t.Fatalf("expected daily quota usage reset, got %d", dailyUsed)
|
||||
}
|
||||
quotaDisabled := mustQueryInt(t, r, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`)
|
||||
if quotaDisabled != 0 {
|
||||
t.Fatalf("expected quota disabled flag cleared, got %d", quotaDisabled)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,940 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultMetricsRangeMs = int64(60 * 60 * 1000) // 1h
|
||||
maxMetricsRangeMs = int64(24 * 60 * 60 * 1000) // 24h
|
||||
)
|
||||
|
||||
func (h *Handler) resolveServiceMonitorLimits() monitoring.ServiceMonitorLimits {
|
||||
defaults := monitoring.DefaultServiceMonitorLimits()
|
||||
if h == nil || h.repo == nil {
|
||||
return defaults
|
||||
}
|
||||
cfg, err := h.repo.GetConfigsByNames([]string{
|
||||
monitoring.ConfigServiceMonitorCheckerScanIntervalSec,
|
||||
monitoring.ConfigServiceMonitorWorkerLimit,
|
||||
monitoring.ConfigServiceMonitorMinIntervalSec,
|
||||
monitoring.ConfigServiceMonitorDefaultIntervalSec,
|
||||
monitoring.ConfigServiceMonitorMinTimeoutSec,
|
||||
monitoring.ConfigServiceMonitorDefaultTimeoutSec,
|
||||
monitoring.ConfigServiceMonitorMaxTimeoutSec,
|
||||
})
|
||||
if err != nil {
|
||||
return defaults
|
||||
}
|
||||
return monitoring.ServiceMonitorLimitsFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
func (h *Handler) monitorNodeMetricsHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
path := r.URL.Path
|
||||
prefix := "/api/v1/monitor/nodes/"
|
||||
if !strings.HasPrefix(path, prefix) {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的路径"))
|
||||
return
|
||||
}
|
||||
|
||||
rest := strings.TrimPrefix(path, prefix)
|
||||
if strings.HasSuffix(rest, "/metrics/latest") {
|
||||
h.handleNodeMetricsLatest(w, r, strings.TrimSuffix(rest, "/metrics/latest"))
|
||||
return
|
||||
}
|
||||
if strings.HasSuffix(rest, "/metrics") {
|
||||
h.handleNodeMetrics(w, r, strings.TrimSuffix(rest, "/metrics"))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.ErrDefault("无效的路径"))
|
||||
}
|
||||
|
||||
type monitorNodeListItem struct {
|
||||
ID int64 `json:"id"`
|
||||
Inx int `json:"inx"`
|
||||
Name string `json:"name"`
|
||||
Status int `json:"status"`
|
||||
Version string `json:"version"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
}
|
||||
|
||||
func (h *Handler) monitorNodeListHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
nodes, err := h.repo.ListMonitorNodes()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
items := make([]monitorNodeListItem, 0, len(nodes))
|
||||
for _, n := range nodes {
|
||||
updated := int64(0)
|
||||
if n.UpdatedTime.Valid {
|
||||
updated = n.UpdatedTime.Int64
|
||||
}
|
||||
items = append(items, monitorNodeListItem{
|
||||
ID: n.ID,
|
||||
Inx: n.Inx,
|
||||
Name: n.Name,
|
||||
Status: n.Status,
|
||||
Version: n.Version.String,
|
||||
UpdatedTime: updated,
|
||||
})
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
type monitorTunnelListItem struct {
|
||||
ID int64 `json:"id"`
|
||||
Inx int `json:"inx"`
|
||||
Name string `json:"name"`
|
||||
Status int `json:"status"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
}
|
||||
|
||||
func (h *Handler) monitorTunnelListHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
tunnels, err := h.repo.ListMonitorTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
items := make([]monitorTunnelListItem, 0, len(tunnels))
|
||||
for _, t := range tunnels {
|
||||
items = append(items, monitorTunnelListItem{
|
||||
ID: t.ID,
|
||||
Inx: t.Inx,
|
||||
Name: t.Name,
|
||||
Status: t.Status,
|
||||
UpdatedTime: t.UpdatedTime,
|
||||
})
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
func (h *Handler) handleNodeMetrics(w http.ResponseWriter, r *http.Request, nodeIDStr string) {
|
||||
nodeID, err := strconv.ParseInt(nodeIDStr, 10, 64)
|
||||
if err != nil || nodeID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的节点ID"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
startMs := now - defaultMetricsRangeMs
|
||||
endMs := now
|
||||
|
||||
if s := r.URL.Query().Get("start"); s != "" {
|
||||
if v, err := strconv.ParseInt(s, 10, 64); err == nil {
|
||||
startMs = v
|
||||
}
|
||||
}
|
||||
if e := r.URL.Query().Get("end"); e != "" {
|
||||
if v, err := strconv.ParseInt(e, 10, 64); err == nil {
|
||||
endMs = v
|
||||
}
|
||||
}
|
||||
if startMs <= 0 || endMs <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的时间范围"))
|
||||
return
|
||||
}
|
||||
if endMs < startMs {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的时间范围"))
|
||||
return
|
||||
}
|
||||
if endMs-startMs > maxMetricsRangeMs {
|
||||
response.WriteJSON(w, response.ErrDefault("时间范围过大"))
|
||||
return
|
||||
}
|
||||
|
||||
metrics, err := h.repo.GetNodeMetrics(nodeID, startMs, endMs)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(metrics))
|
||||
}
|
||||
|
||||
func (h *Handler) handleNodeMetricsLatest(w http.ResponseWriter, _ *http.Request, nodeIDStr string) {
|
||||
nodeID, err := strconv.ParseInt(nodeIDStr, 10, 64)
|
||||
if err != nil || nodeID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的节点ID"))
|
||||
return
|
||||
}
|
||||
|
||||
metric, err := h.repo.GetLatestNodeMetric(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if metric == nil {
|
||||
response.WriteJSON(w, response.OK(nil))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(metric))
|
||||
}
|
||||
|
||||
func (h *Handler) monitorTunnelQualityHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
// Try in-memory cache first
|
||||
if h.qualityProber != nil {
|
||||
items := h.qualityProber.GetAll()
|
||||
if len(items) > 0 {
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to database (latest per tunnel)
|
||||
qualities, err := h.repo.GetLatestTunnelQualities()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
snapshots := make([]tunnelQualitySnapshot, 0, len(qualities))
|
||||
for _, q := range qualities {
|
||||
snapshots = append(snapshots, tunnelQualitySnapshot{
|
||||
TunnelID: q.TunnelID,
|
||||
EntryToExitLatency: q.EntryToExitLatency,
|
||||
ExitToBingLatency: q.ExitToBingLatency,
|
||||
EntryToExitLoss: q.EntryToExitLoss,
|
||||
ExitToBingLoss: q.ExitToBingLoss,
|
||||
Success: q.Success == 1,
|
||||
ErrorMessage: q.ErrorMessage,
|
||||
Timestamp: q.Timestamp,
|
||||
ChainDetails: q.ChainDetails,
|
||||
})
|
||||
}
|
||||
response.WriteJSON(w, response.OK(snapshots))
|
||||
}
|
||||
|
||||
// monitorTunnelQualityHistory returns quality probe history for charting.
|
||||
// GET /api/v1/monitor/tunnels/{id}/quality?start=...&end=...
|
||||
// Mirrors monitorTunnelMetrics / monitorServiceResultsHandler pattern.
|
||||
func (h *Handler) monitorTunnelQualityHistory(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
tunnelIDStr := extractPathParam(r.URL.Path, "/api/v1/monitor/tunnels/", "/quality")
|
||||
tunnelID, err := strconv.ParseInt(tunnelIDStr, 10, 64)
|
||||
if err != nil || tunnelID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的隧道ID"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
startMs := now - defaultMetricsRangeMs
|
||||
endMs := now
|
||||
|
||||
if s := r.URL.Query().Get("start"); s != "" {
|
||||
if v, err := strconv.ParseInt(s, 10, 64); err == nil {
|
||||
startMs = v
|
||||
}
|
||||
}
|
||||
if e := r.URL.Query().Get("end"); e != "" {
|
||||
if v, err := strconv.ParseInt(e, 10, 64); err == nil {
|
||||
endMs = v
|
||||
}
|
||||
}
|
||||
if startMs <= 0 || endMs <= 0 || endMs < startMs {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的时间范围"))
|
||||
return
|
||||
}
|
||||
if endMs-startMs > maxMetricsRangeMs {
|
||||
response.WriteJSON(w, response.ErrDefault("时间范围过大"))
|
||||
return
|
||||
}
|
||||
|
||||
results, err := h.repo.GetTunnelQualityHistory(tunnelID, startMs, endMs)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(results))
|
||||
}
|
||||
|
||||
func (h *Handler) monitorTunnelMetrics(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
path := r.URL.Path
|
||||
prefix := "/api/v1/monitor/tunnels/"
|
||||
if !strings.HasPrefix(path, prefix) {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的路径"))
|
||||
return
|
||||
}
|
||||
|
||||
rest := strings.TrimPrefix(path, prefix)
|
||||
|
||||
// Route: /api/v1/monitor/tunnels/{id}/quality
|
||||
if strings.HasSuffix(rest, "/quality") {
|
||||
h.monitorTunnelQualityHistory(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// Route: /api/v1/monitor/tunnels/{id}/metrics (original)
|
||||
tunnelIDStr := extractPathParam(path, prefix, "/metrics")
|
||||
tunnelID, err := strconv.ParseInt(tunnelIDStr, 10, 64)
|
||||
if err != nil || tunnelID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的隧道ID"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
startMs := now - defaultMetricsRangeMs
|
||||
endMs := now
|
||||
|
||||
if s := r.URL.Query().Get("start"); s != "" {
|
||||
if v, err := strconv.ParseInt(s, 10, 64); err == nil {
|
||||
startMs = v
|
||||
}
|
||||
}
|
||||
if e := r.URL.Query().Get("end"); e != "" {
|
||||
if v, err := strconv.ParseInt(e, 10, 64); err == nil {
|
||||
endMs = v
|
||||
}
|
||||
}
|
||||
if startMs <= 0 || endMs <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的时间范围"))
|
||||
return
|
||||
}
|
||||
if endMs < startMs {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的时间范围"))
|
||||
return
|
||||
}
|
||||
if endMs-startMs > maxMetricsRangeMs {
|
||||
response.WriteJSON(w, response.ErrDefault("时间范围过大"))
|
||||
return
|
||||
}
|
||||
|
||||
metrics, err := h.repo.GetTunnelMetricsAggregated(tunnelID, startMs, endMs)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(metrics))
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceListHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
monitors, err := h.repo.ListServiceMonitors()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(monitors))
|
||||
}
|
||||
|
||||
type createServiceMonitorRequest struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Target string `json:"target"`
|
||||
IntervalSec int `json:"intervalSec"`
|
||||
TimeoutSec int `json:"timeoutSec"`
|
||||
NodeID int64 `json:"nodeId"`
|
||||
Enabled *int `json:"enabled"`
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
var req createServiceMonitorRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
name := strings.TrimSpace(req.Name)
|
||||
if name == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("名称不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
monitorType := strings.ToLower(strings.TrimSpace(req.Type))
|
||||
if monitorType != "tcp" && monitorType != "icmp" {
|
||||
response.WriteJSON(w, response.ErrDefault("类型必须是 tcp 或 icmp"))
|
||||
return
|
||||
}
|
||||
|
||||
target := strings.TrimSpace(req.Target)
|
||||
if target == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("目标地址不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
limits := h.resolveServiceMonitorLimits()
|
||||
|
||||
intervalSec := req.IntervalSec
|
||||
if intervalSec <= 0 {
|
||||
intervalSec = limits.DefaultIntervalSec
|
||||
}
|
||||
if intervalSec < limits.MinIntervalSec {
|
||||
intervalSec = limits.MinIntervalSec
|
||||
}
|
||||
|
||||
timeoutSec := req.TimeoutSec
|
||||
if timeoutSec <= 0 {
|
||||
timeoutSec = limits.DefaultTimeoutSec
|
||||
}
|
||||
if timeoutSec < limits.MinTimeoutSec {
|
||||
timeoutSec = limits.MinTimeoutSec
|
||||
}
|
||||
if timeoutSec > limits.MaxTimeoutSec {
|
||||
timeoutSec = limits.MaxTimeoutSec
|
||||
}
|
||||
|
||||
enabled := 1
|
||||
if req.Enabled != nil {
|
||||
if *req.Enabled == 0 || *req.Enabled == 1 {
|
||||
enabled = *req.Enabled
|
||||
}
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if req.NodeID > 0 {
|
||||
n, err := h.repo.GetNodeByID(req.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if n == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("节点不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
m := &model.ServiceMonitor{
|
||||
Name: name,
|
||||
Type: monitorType,
|
||||
Target: target,
|
||||
IntervalSec: intervalSec,
|
||||
TimeoutSec: timeoutSec,
|
||||
NodeID: req.NodeID,
|
||||
Enabled: enabled,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}
|
||||
if m.Type == "icmp" && m.NodeID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("ICMP 监控必须选择执行节点"))
|
||||
return
|
||||
}
|
||||
// enabled is already normalized above.
|
||||
|
||||
if err := h.repo.CreateServiceMonitor(m); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(m))
|
||||
}
|
||||
|
||||
type updateServiceMonitorRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Target string `json:"target"`
|
||||
IntervalSec int `json:"intervalSec"`
|
||||
TimeoutSec int `json:"timeoutSec"`
|
||||
NodeID *int64 `json:"nodeId"`
|
||||
Enabled *int `json:"enabled"`
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
var req updateServiceMonitorRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的监控ID"))
|
||||
return
|
||||
}
|
||||
|
||||
existing, err := h.repo.GetServiceMonitor(req.ID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if existing == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("监控不存在"))
|
||||
return
|
||||
}
|
||||
|
||||
name := strings.TrimSpace(req.Name)
|
||||
if name != "" {
|
||||
existing.Name = name
|
||||
}
|
||||
|
||||
monitorType := strings.ToLower(strings.TrimSpace(req.Type))
|
||||
if monitorType == "tcp" || monitorType == "icmp" {
|
||||
existing.Type = monitorType
|
||||
}
|
||||
|
||||
target := strings.TrimSpace(req.Target)
|
||||
if target != "" {
|
||||
existing.Target = target
|
||||
}
|
||||
|
||||
limits := h.resolveServiceMonitorLimits()
|
||||
|
||||
if req.IntervalSec > 0 {
|
||||
intervalSec := req.IntervalSec
|
||||
if intervalSec < limits.MinIntervalSec {
|
||||
intervalSec = limits.MinIntervalSec
|
||||
}
|
||||
existing.IntervalSec = intervalSec
|
||||
}
|
||||
if req.TimeoutSec > 0 {
|
||||
timeoutSec := req.TimeoutSec
|
||||
if timeoutSec < limits.MinTimeoutSec {
|
||||
timeoutSec = limits.MinTimeoutSec
|
||||
}
|
||||
if timeoutSec > limits.MaxTimeoutSec {
|
||||
timeoutSec = limits.MaxTimeoutSec
|
||||
}
|
||||
existing.TimeoutSec = timeoutSec
|
||||
}
|
||||
|
||||
if req.NodeID != nil {
|
||||
existing.NodeID = *req.NodeID
|
||||
}
|
||||
if req.Enabled != nil {
|
||||
if *req.Enabled == 0 || *req.Enabled == 1 {
|
||||
existing.Enabled = *req.Enabled
|
||||
}
|
||||
}
|
||||
|
||||
existing.UpdatedTime = time.Now().UnixMilli()
|
||||
if existing.Type == "icmp" && existing.NodeID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("ICMP 监控必须选择执行节点"))
|
||||
return
|
||||
}
|
||||
if existing.NodeID > 0 {
|
||||
n, err := h.repo.GetNodeByID(existing.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if n == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("节点不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateServiceMonitor(existing); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(existing))
|
||||
}
|
||||
|
||||
type deleteServiceMonitorRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
var req deleteServiceMonitorRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的监控ID"))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.DeleteServiceMonitor(req.ID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceRun(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
if h.healthCheck == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("监控服务不可用"))
|
||||
return
|
||||
}
|
||||
|
||||
var req deleteServiceMonitorRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的监控ID"))
|
||||
return
|
||||
}
|
||||
|
||||
m, err := h.repo.GetServiceMonitor(req.ID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if m == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("监控不存在"))
|
||||
return
|
||||
}
|
||||
|
||||
res, err := h.healthCheck.RunOnce(m)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.InsertServiceMonitorResult(res); err != nil {
|
||||
log.Printf("monitoring write failed op=service_monitor_result.manual_insert monitor_id=%d err=%v", res.MonitorID, err)
|
||||
}
|
||||
response.WriteJSON(w, response.OK(res))
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceResultsHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
monitorIDStr := extractPathParam(r.URL.Path, "/api/v1/monitor/services/", "/results")
|
||||
monitorID, err := strconv.ParseInt(monitorIDStr, 10, 64)
|
||||
if err != nil || monitorID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的监控ID"))
|
||||
return
|
||||
}
|
||||
|
||||
// If start/end time range is provided, use time-based query (mirrors node metrics / tunnel quality pattern).
|
||||
startStr := r.URL.Query().Get("start")
|
||||
endStr := r.URL.Query().Get("end")
|
||||
if startStr != "" && endStr != "" {
|
||||
startMs, err1 := strconv.ParseInt(startStr, 10, 64)
|
||||
endMs, err2 := strconv.ParseInt(endStr, 10, 64)
|
||||
if err1 != nil || err2 != nil || startMs <= 0 || endMs <= 0 || endMs < startMs {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的时间范围"))
|
||||
return
|
||||
}
|
||||
if endMs-startMs > maxMetricsRangeMs {
|
||||
response.WriteJSON(w, response.ErrDefault("时间范围过大"))
|
||||
return
|
||||
}
|
||||
results, err := h.repo.GetServiceMonitorResultsByTimeRange(monitorID, startMs, endMs)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(results))
|
||||
return
|
||||
}
|
||||
|
||||
// Fallback: count-based limit query (backward compat).
|
||||
limit := 100
|
||||
if l := r.URL.Query().Get("limit"); l != "" {
|
||||
if v, err := strconv.Atoi(l); err == nil && v > 0 && v <= 1000 {
|
||||
limit = v
|
||||
}
|
||||
}
|
||||
|
||||
results, err := h.repo.GetServiceMonitorResults(monitorID, limit)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(results))
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceLatestResultsHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
// Try in-memory cache first (updated every 1s)
|
||||
if h.healthCheck != nil {
|
||||
cached := h.healthCheck.GetLatestCached()
|
||||
if len(cached) > 0 {
|
||||
response.WriteJSON(w, response.OK(cached))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback to database
|
||||
results, err := h.repo.GetLatestServiceMonitorResults()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(results))
|
||||
}
|
||||
|
||||
func (h *Handler) monitorServiceLimitsHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureMonitoringAccess(w, r) {
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(h.resolveServiceMonitorLimits()))
|
||||
}
|
||||
|
||||
func extractPathParam(path, prefix, suffix string) string {
|
||||
if !strings.HasPrefix(path, prefix) {
|
||||
return ""
|
||||
}
|
||||
rest := strings.TrimPrefix(path, prefix)
|
||||
if suffix != "" {
|
||||
rest = strings.TrimSuffix(rest, suffix)
|
||||
}
|
||||
return rest
|
||||
}
|
||||
|
||||
type monitorAccessData struct {
|
||||
Allowed bool `json:"allowed"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
}
|
||||
|
||||
// monitorAccessHandler is a lightweight capability check for frontend navigation.
|
||||
// It does NOT replace authorization on the actual monitoring endpoints.
|
||||
func (h *Handler) monitorAccessHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
userID, roleID, err := userRoleFromRequest(r)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
|
||||
return
|
||||
}
|
||||
if roleID == 0 {
|
||||
response.WriteJSON(w, response.OK(monitorAccessData{Allowed: true}))
|
||||
return
|
||||
}
|
||||
|
||||
allowed, err := h.repo.HasMonitorPermission(userID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
data := monitorAccessData{Allowed: allowed}
|
||||
if !allowed {
|
||||
data.Reason = "need_admin_grant"
|
||||
}
|
||||
response.WriteJSON(w, response.OK(data))
|
||||
}
|
||||
|
||||
func (h *Handler) ensureAdminAccess(w http.ResponseWriter, r *http.Request) bool {
|
||||
_, roleID, err := userRoleFromRequest(r)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
|
||||
return false
|
||||
}
|
||||
if roleID != 0 {
|
||||
response.WriteJSON(w, response.Err(403, "权限不足,仅管理员可操作"))
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (h *Handler) ensureMonitoringAccess(w http.ResponseWriter, r *http.Request) bool {
|
||||
userID, roleID, err := userRoleFromRequest(r)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(401, "未登录或token已过期"))
|
||||
return false
|
||||
}
|
||||
if roleID == 0 {
|
||||
return true
|
||||
}
|
||||
allowed, err := h.repo.HasMonitorPermission(userID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return false
|
||||
}
|
||||
if !allowed {
|
||||
response.WriteJSON(w, response.Err(403, "权限不足:当前账户非管理员,且未被授予监控权限。请联系管理员在用户管理中授权监控权限。"))
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (h *Handler) monitorPermissionList(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureAdminAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
items, err := h.repo.ListMonitorPermissions()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
type monitorPermissionMutationRequest struct {
|
||||
UserID int64 `json:"userId"`
|
||||
}
|
||||
|
||||
func (h *Handler) monitorPermissionAssign(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureAdminAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
var req monitorPermissionMutationRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.UserID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的用户ID"))
|
||||
return
|
||||
}
|
||||
|
||||
u, err := h.repo.GetUserByID(req.UserID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if u == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("用户不存在"))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.InsertMonitorPermission(req.UserID, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) monitorPermissionRemove(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.ensureAdminAccess(w, r) {
|
||||
return
|
||||
}
|
||||
|
||||
var req monitorPermissionMutationRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.UserID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("无效的用户ID"))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.DeleteMonitorPermission(req.UserID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,39 @@
|
||||
package handler
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestBuildForwardPortEntriesWithPreservedInIP(t *testing.T) {
|
||||
entryNodeIDs := []int64{10, 20, 30}
|
||||
oldPorts := []forwardPortRecord{
|
||||
{NodeID: 10, Port: 10001, InIP: ""},
|
||||
{NodeID: 10, Port: 10002, InIP: "10.0.0.10"},
|
||||
{NodeID: 20, Port: 10003, InIP: "10.0.0.20"},
|
||||
}
|
||||
|
||||
entries := buildForwardPortEntriesWithPreservedInIP(entryNodeIDs, oldPorts, 18080)
|
||||
if len(entries) != 3 {
|
||||
t.Fatalf("expected 3 entries, got %d", len(entries))
|
||||
}
|
||||
|
||||
if entries[0].NodeID != 10 || entries[0].Port != 18080 || entries[0].InIP != "10.0.0.10" {
|
||||
t.Fatalf("unexpected first entry: %+v", entries[0])
|
||||
}
|
||||
if entries[1].NodeID != 20 || entries[1].Port != 18080 || entries[1].InIP != "10.0.0.20" {
|
||||
t.Fatalf("unexpected second entry: %+v", entries[1])
|
||||
}
|
||||
if entries[2].NodeID != 30 || entries[2].Port != 18080 || entries[2].InIP != "" {
|
||||
t.Fatalf("unexpected third entry: %+v", entries[2])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardPortEntriesWithPreservedInIP_EmptyOldPorts(t *testing.T) {
|
||||
entryNodeIDs := []int64{99}
|
||||
entries := buildForwardPortEntriesWithPreservedInIP(entryNodeIDs, nil, 17000)
|
||||
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("expected 1 entry, got %d", len(entries))
|
||||
}
|
||||
if entries[0].NodeID != 99 || entries[0].Port != 17000 || entries[0].InIP != "" {
|
||||
t.Fatalf("unexpected entry: %+v", entries[0])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestReconstructTunnelState_PreservesConnectIP(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "reconstruct-connect-ip.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'reconstruct-tunnel', 1.0, 2, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
insertNode := func(id int64, name, ip string) {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, id, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
insertNode(101, "entry", "10.90.0.10")
|
||||
insertNode(102, "middle", "10.90.0.20")
|
||||
insertNode(103, "exit", "10.90.0.30")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(1, '1', 101, 30001, 'round', 1, 'tls')
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert entry chain: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol, connect_ip)
|
||||
VALUES(1, '2', 102, 30002, 'round', 1, 'tls', '10.99.9.22')
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert middle chain: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol, connect_ip)
|
||||
VALUES(1, '3', 103, 30003, 'round', 1, 'tls', '10.99.9.33')
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert exit chain: %v", err)
|
||||
}
|
||||
|
||||
state, err := h.reconstructTunnelState(1)
|
||||
if err != nil {
|
||||
t.Fatalf("reconstructTunnelState: %v", err)
|
||||
}
|
||||
|
||||
if len(state.ChainHops) != 1 || len(state.ChainHops[0]) != 1 {
|
||||
t.Fatalf("unexpected chain hops: %+v", state.ChainHops)
|
||||
}
|
||||
if got := state.ChainHops[0][0].ConnectIP; got != "10.99.9.22" {
|
||||
t.Fatalf("expected middle connectIp 10.99.9.22, got %q", got)
|
||||
}
|
||||
|
||||
if len(state.OutNodes) != 1 {
|
||||
t.Fatalf("unexpected out nodes: %+v", state.OutNodes)
|
||||
}
|
||||
if got := state.OutNodes[0].ConnectIP; got != "10.99.9.33" {
|
||||
t.Fatalf("expected exit connectIp 10.99.9.33, got %q", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,655 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
const tunnelDeletePreviewSampleLimit = 5
|
||||
|
||||
const (
|
||||
tunnelDeleteActionReplace = "replace"
|
||||
tunnelDeleteActionDeleteForwards = "delete_forwards"
|
||||
)
|
||||
|
||||
var (
|
||||
errInvalidTunnelDeleteTarget = errors.New("invalid tunnel delete target")
|
||||
)
|
||||
|
||||
type tunnelDeleteForwardPreviewItem struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
UserID int64 `json:"userId"`
|
||||
UserName string `json:"userName"`
|
||||
InPort int `json:"inPort"`
|
||||
}
|
||||
|
||||
type tunnelDeletePreviewData struct {
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
TunnelName string `json:"tunnelName"`
|
||||
ForwardCount int `json:"forwardCount"`
|
||||
SampleForwards []tunnelDeleteForwardPreviewItem `json:"sampleForwards"`
|
||||
}
|
||||
|
||||
type tunnelBatchDeletePreviewData struct {
|
||||
TunnelCount int `json:"tunnelCount"`
|
||||
TotalForwardCount int `json:"totalForwardCount"`
|
||||
Items []tunnelDeletePreviewData `json:"items"`
|
||||
}
|
||||
|
||||
type tunnelDeleteWithForwardsRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
Action string `json:"action"`
|
||||
TargetTunnelID int64 `json:"targetTunnelId"`
|
||||
}
|
||||
|
||||
type tunnelBatchDeleteWithForwardsRequest struct {
|
||||
IDs []int64 `json:"ids"`
|
||||
Action string `json:"action"`
|
||||
TargetTunnelID int64 `json:"targetTunnelId"`
|
||||
}
|
||||
|
||||
type tunnelDeleteWithForwardsResult struct {
|
||||
ForwardCount int `json:"forwardCount"`
|
||||
MigratedCount int `json:"migratedCount"`
|
||||
DeletedForwardCount int `json:"deletedForwardCount"`
|
||||
PortAdjustedCount int `json:"portAdjustedCount"`
|
||||
Warnings []string `json:"warnings,omitempty"`
|
||||
}
|
||||
|
||||
type tunnelBatchDeleteWithForwardsResult struct {
|
||||
SuccessCount int `json:"successCount"`
|
||||
FailCount int `json:"failCount"`
|
||||
Failures []batchFailureDetail `json:"failures,omitempty"`
|
||||
DeletedForwardCount int `json:"deletedForwardCount"`
|
||||
MigratedCount int `json:"migratedCount"`
|
||||
PortAdjustedCount int `json:"portAdjustedCount"`
|
||||
Warnings []string `json:"warnings,omitempty"`
|
||||
}
|
||||
|
||||
type tunnelForwardMigrationPlan struct {
|
||||
forward *forwardRecord
|
||||
oldPorts []forwardPortRecord
|
||||
targetTunnelID int64
|
||||
targetPort int
|
||||
keptNodeIDs []int64
|
||||
removedNodeIDs []int64
|
||||
portAdjusted bool
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelDeletePreview(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
id := idFromBody(r, w)
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
preview, err := h.buildTunnelDeletePreview(id)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "不存在") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(preview))
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelBatchDeletePreview(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
IDs []int64 `json:"ids"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil || len(req.IDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
preview, err := h.buildTunnelBatchDeletePreview(req.IDs)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "不存在") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(preview))
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelDeleteWithForwards(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req tunnelDeleteWithForwardsRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil || req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
action, err := normalizeTunnelDeleteAction(req.Action)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if action == tunnelDeleteActionReplace {
|
||||
if _, _, authErr := userRoleFromRequest(r); authErr != nil {
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
result, failures, err := h.processTunnelDeleteWithForwards(req.ID, action, req.TargetTunnelID)
|
||||
if err != nil {
|
||||
if err == errInvalidTunnelDeleteTarget {
|
||||
response.WriteJSON(w, response.ErrDefault("目标隧道不能为空"))
|
||||
return
|
||||
}
|
||||
if strings.Contains(err.Error(), "目标隧道不能与当前隧道相同") || strings.Contains(err.Error(), "目标隧道不存在") || strings.Contains(err.Error(), "目标隧道已禁用") || strings.Contains(err.Error(), "隧道不存在") {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if len(failures) > 0 {
|
||||
response.WriteJSON(w, response.R{
|
||||
Code: -2,
|
||||
Msg: "部分规则迁移失败",
|
||||
TS: time.Now().UnixMilli(),
|
||||
Data: batchOperationResult{SuccessCount: 0, FailCount: len(failures), Failures: failures},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(result))
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelBatchDeleteWithForwards(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req tunnelBatchDeleteWithForwardsRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil || len(req.IDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
action, err := normalizeTunnelDeleteAction(req.Action)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if action == tunnelDeleteActionReplace {
|
||||
if _, _, authErr := userRoleFromRequest(r); authErr != nil {
|
||||
response.WriteJSON(w, response.Err(401, "无效的token或token已过期"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
normalizedIDs := normalizeTunnelIDs(req.IDs)
|
||||
if len(normalizedIDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
if action == tunnelDeleteActionReplace {
|
||||
if req.TargetTunnelID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("目标隧道不能为空"))
|
||||
return
|
||||
}
|
||||
for _, id := range normalizedIDs {
|
||||
if id == req.TargetTunnelID {
|
||||
response.WriteJSON(w, response.ErrDefault("目标隧道不能包含在删除列表中"))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result := tunnelBatchDeleteWithForwardsResult{}
|
||||
for _, tunnelID := range normalizedIDs {
|
||||
tunnelName, _ := h.repo.GetTunnelName(tunnelID)
|
||||
singleResult, failures, processErr := h.processTunnelDeleteWithForwards(tunnelID, action, req.TargetTunnelID)
|
||||
if processErr != nil {
|
||||
result.FailCount++
|
||||
result.Failures = appendBatchFailure(result.Failures, tunnelID, tunnelName, processErr)
|
||||
continue
|
||||
}
|
||||
if len(failures) > 0 {
|
||||
result.FailCount++
|
||||
result.Failures = appendBatchFailureReason(
|
||||
result.Failures,
|
||||
tunnelID,
|
||||
tunnelName,
|
||||
summarizeTunnelDeleteRuleFailures(failures),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
result.SuccessCount++
|
||||
result.DeletedForwardCount += singleResult.DeletedForwardCount
|
||||
result.MigratedCount += singleResult.MigratedCount
|
||||
result.PortAdjustedCount += singleResult.PortAdjustedCount
|
||||
if len(singleResult.Warnings) > 0 {
|
||||
result.Warnings = append(result.Warnings, singleResult.Warnings...)
|
||||
}
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(result))
|
||||
}
|
||||
|
||||
func (h *Handler) buildTunnelDeletePreview(tunnelID int64) (*tunnelDeletePreviewData, error) {
|
||||
if _, err := h.getTunnelRecord(tunnelID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
tunnelName, err := h.repo.GetTunnelName(tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
forwards, err := h.listForwardsByTunnel(tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
samples := make([]tunnelDeleteForwardPreviewItem, 0, minInt(len(forwards), tunnelDeletePreviewSampleLimit))
|
||||
for i, forward := range forwards {
|
||||
if i >= tunnelDeletePreviewSampleLimit {
|
||||
break
|
||||
}
|
||||
ports, portsErr := h.listForwardPorts(forward.ID)
|
||||
if portsErr != nil {
|
||||
return nil, portsErr
|
||||
}
|
||||
inPort := 0
|
||||
if len(ports) > 0 {
|
||||
inPort = ports[0].Port
|
||||
}
|
||||
samples = append(samples, tunnelDeleteForwardPreviewItem{
|
||||
ID: forward.ID,
|
||||
Name: forward.Name,
|
||||
UserID: forward.UserID,
|
||||
UserName: forward.UserName,
|
||||
InPort: inPort,
|
||||
})
|
||||
}
|
||||
|
||||
return &tunnelDeletePreviewData{
|
||||
TunnelID: tunnelID,
|
||||
TunnelName: tunnelName,
|
||||
ForwardCount: len(forwards),
|
||||
SampleForwards: samples,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) buildTunnelBatchDeletePreview(ids []int64) (*tunnelBatchDeletePreviewData, error) {
|
||||
normalizedIDs := normalizeTunnelIDs(ids)
|
||||
items := make([]tunnelDeletePreviewData, 0, len(normalizedIDs))
|
||||
totalForwardCount := 0
|
||||
for _, id := range normalizedIDs {
|
||||
preview, err := h.buildTunnelDeletePreview(id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, *preview)
|
||||
totalForwardCount += preview.ForwardCount
|
||||
}
|
||||
return &tunnelBatchDeletePreviewData{
|
||||
TunnelCount: len(items),
|
||||
TotalForwardCount: totalForwardCount,
|
||||
Items: items,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func normalizeTunnelDeleteAction(action string) (string, error) {
|
||||
normalized := strings.TrimSpace(action)
|
||||
if normalized == "" {
|
||||
return tunnelDeleteActionDeleteForwards, nil
|
||||
}
|
||||
if normalized != tunnelDeleteActionReplace && normalized != tunnelDeleteActionDeleteForwards {
|
||||
return "", errors.New("invalid tunnel delete action")
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func normalizeTunnelIDs(ids []int64) []int64 {
|
||||
seen := make(map[int64]struct{}, len(ids))
|
||||
out := make([]int64, 0, len(ids))
|
||||
for _, id := range ids {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[id]; exists {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
out = append(out, id)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func summarizeTunnelDeleteRuleFailures(failures []batchFailureDetail) string {
|
||||
if len(failures) == 0 {
|
||||
return "未知错误"
|
||||
}
|
||||
parts := make([]string, 0, minInt(len(failures), 3))
|
||||
for i, failure := range failures {
|
||||
if i >= 3 {
|
||||
break
|
||||
}
|
||||
name := strings.TrimSpace(failure.Name)
|
||||
if name == "" {
|
||||
name = fmt.Sprintf("规则 #%d", failure.ID)
|
||||
}
|
||||
parts = append(parts, fmt.Sprintf("%s: %s", name, strings.TrimSpace(failure.Reason)))
|
||||
}
|
||||
if len(failures) > 3 {
|
||||
parts = append(parts, fmt.Sprintf("另有 %d 条规则失败", len(failures)-3))
|
||||
}
|
||||
return strings.Join(parts, ";")
|
||||
}
|
||||
|
||||
func (h *Handler) processTunnelDeleteWithForwards(tunnelID int64, action string, targetTunnelID int64) (tunnelDeleteWithForwardsResult, []batchFailureDetail, error) {
|
||||
preview, err := h.buildTunnelDeletePreview(tunnelID)
|
||||
if err != nil {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, err
|
||||
}
|
||||
|
||||
result := tunnelDeleteWithForwardsResult{ForwardCount: preview.ForwardCount}
|
||||
if preview.ForwardCount == 0 {
|
||||
if err := h.deleteTunnelAndCleanup(tunnelID); err != nil {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, err
|
||||
}
|
||||
return result, nil, nil
|
||||
}
|
||||
|
||||
if action == tunnelDeleteActionDeleteForwards {
|
||||
result.DeletedForwardCount = preview.ForwardCount
|
||||
if err := h.deleteTunnelAndCleanup(tunnelID); err != nil {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, err
|
||||
}
|
||||
return result, nil, nil
|
||||
}
|
||||
|
||||
if targetTunnelID <= 0 {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, errInvalidTunnelDeleteTarget
|
||||
}
|
||||
if targetTunnelID == tunnelID {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, errors.New("目标隧道不能与当前隧道相同")
|
||||
}
|
||||
return h.processTunnelDeleteReplaceAction(tunnelID, targetTunnelID, result)
|
||||
}
|
||||
|
||||
func (h *Handler) processTunnelDeleteReplaceAction(tunnelID, targetTunnelID int64, result tunnelDeleteWithForwardsResult) (tunnelDeleteWithForwardsResult, []batchFailureDetail, error) {
|
||||
targetTunnel, err := h.getTunnelRecord(targetTunnelID)
|
||||
if err != nil {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, errors.New("目标隧道不存在")
|
||||
}
|
||||
if targetTunnel.Status != 1 {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, errors.New("目标隧道已禁用")
|
||||
}
|
||||
|
||||
plans, failures, err := h.planTunnelDeleteForwardMigrations(tunnelID, targetTunnelID)
|
||||
if err != nil {
|
||||
return tunnelDeleteWithForwardsResult{}, nil, err
|
||||
}
|
||||
if len(failures) > 0 {
|
||||
return tunnelDeleteWithForwardsResult{}, failures, nil
|
||||
}
|
||||
|
||||
portAdjustedCount := 0
|
||||
warnings, execErr, execFailure := h.executeTunnelDeleteForwardMigrations(plans)
|
||||
for _, plan := range plans {
|
||||
if plan.portAdjusted {
|
||||
portAdjustedCount++
|
||||
}
|
||||
}
|
||||
if execErr != nil {
|
||||
failures = append(failures, execFailure)
|
||||
return tunnelDeleteWithForwardsResult{}, failures, nil
|
||||
}
|
||||
|
||||
if err := h.deleteTunnelAndCleanup(tunnelID); err != nil {
|
||||
h.rollbackTunnelForwardMigrationPlans(plans)
|
||||
_ = h.redeployTunnelAndForwards(tunnelID)
|
||||
return tunnelDeleteWithForwardsResult{}, nil, err
|
||||
}
|
||||
|
||||
result.MigratedCount = len(plans)
|
||||
result.PortAdjustedCount = portAdjustedCount
|
||||
if len(warnings) > 0 {
|
||||
result.Warnings = warnings
|
||||
}
|
||||
return result, nil, nil
|
||||
}
|
||||
|
||||
func (h *Handler) planTunnelDeleteForwardMigrations(sourceTunnelID, targetTunnelID int64) ([]tunnelForwardMigrationPlan, []batchFailureDetail, error) {
|
||||
forwards, err := h.listForwardsByTunnel(sourceTunnelID)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
entryNodes, err := h.tunnelEntryNodeIDs(targetTunnelID)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if len(entryNodes) == 0 {
|
||||
return nil, nil, errors.New("目标隧道缺少入口节点")
|
||||
}
|
||||
|
||||
plans := make([]tunnelForwardMigrationPlan, 0, len(forwards))
|
||||
failures := make([]batchFailureDetail, 0)
|
||||
reservedPorts := make(map[int64]map[int]bool)
|
||||
|
||||
for _, forward := range forwards {
|
||||
plan, planErr := h.planSingleTunnelDeleteForwardMigration(&forward, targetTunnelID, entryNodes, reservedPorts)
|
||||
if planErr != nil {
|
||||
failures = appendBatchFailure(failures, forward.ID, forward.Name, planErr)
|
||||
continue
|
||||
}
|
||||
plans = append(plans, plan)
|
||||
}
|
||||
|
||||
return plans, failures, nil
|
||||
}
|
||||
|
||||
func (h *Handler) planSingleTunnelDeleteForwardMigration(forward *forwardRecord, targetTunnelID int64, targetEntryNodes []int64, reservedPorts map[int64]map[int]bool) (tunnelForwardMigrationPlan, error) {
|
||||
if forward == nil {
|
||||
return tunnelForwardMigrationPlan{}, errors.New("转发不存在")
|
||||
}
|
||||
|
||||
oldPorts, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return tunnelForwardMigrationPlan{}, err
|
||||
}
|
||||
if len(oldPorts) == 0 {
|
||||
return tunnelForwardMigrationPlan{}, errors.New("转发入口端口不存在")
|
||||
}
|
||||
|
||||
minPort := h.repo.GetMinForwardPort(forward.ID)
|
||||
targetPort := 0
|
||||
if minPort.Valid {
|
||||
targetPort = int(minPort.Int64)
|
||||
}
|
||||
if targetPort <= 0 {
|
||||
targetPort = h.pickTunnelPort(targetTunnelID)
|
||||
}
|
||||
if targetPort <= 0 {
|
||||
targetPort = 10000
|
||||
}
|
||||
|
||||
hasCustomInIP := false
|
||||
for _, oldPort := range oldPorts {
|
||||
if strings.TrimSpace(oldPort.InIP) != "" {
|
||||
hasCustomInIP = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if hasCustomInIP && len(targetEntryNodes) > 1 {
|
||||
return tunnelForwardMigrationPlan{}, errors.New("多入口隧道的转发不支持保留自定义监听IP,请先手动调整该规则")
|
||||
}
|
||||
|
||||
for _, nodeID := range targetEntryNodes {
|
||||
node, nodeErr := h.getNodeRecord(nodeID)
|
||||
if nodeErr != nil {
|
||||
return tunnelForwardMigrationPlan{}, nodeErr
|
||||
}
|
||||
if err := validateRemoteNodePort(node, targetPort); err != nil {
|
||||
return tunnelForwardMigrationPlan{}, err
|
||||
}
|
||||
if err := validateLocalNodePort(node, targetPort); err != nil {
|
||||
return tunnelForwardMigrationPlan{}, err
|
||||
}
|
||||
if err := h.validateForwardPortAvailability(node, targetPort, forward.ID); err != nil {
|
||||
return tunnelForwardMigrationPlan{}, err
|
||||
}
|
||||
if reservedOnNode, ok := reservedPorts[nodeID]; ok && reservedOnNode[targetPort] {
|
||||
return tunnelForwardMigrationPlan{}, fmt.Errorf("目标隧道入口节点端口 %d 已被本次迁移中的其他规则占用", targetPort)
|
||||
}
|
||||
}
|
||||
|
||||
for _, nodeID := range targetEntryNodes {
|
||||
reservedOnNode := reservedPorts[nodeID]
|
||||
if reservedOnNode == nil {
|
||||
reservedOnNode = make(map[int]bool)
|
||||
reservedPorts[nodeID] = reservedOnNode
|
||||
}
|
||||
reservedOnNode[targetPort] = true
|
||||
}
|
||||
|
||||
oldNodeIDs := forwardPortNodeIDs(oldPorts)
|
||||
newNodeIDs := uniqueInt64s(targetEntryNodes)
|
||||
removedNodeIDs := diffInt64s(oldNodeIDs, newNodeIDs)
|
||||
keptNodeIDs := diffInt64s(oldNodeIDs, removedNodeIDs)
|
||||
|
||||
previousPort := 0
|
||||
if len(oldPorts) > 0 {
|
||||
previousPort = oldPorts[0].Port
|
||||
}
|
||||
|
||||
return tunnelForwardMigrationPlan{
|
||||
forward: forward,
|
||||
oldPorts: oldPorts,
|
||||
targetTunnelID: targetTunnelID,
|
||||
targetPort: targetPort,
|
||||
keptNodeIDs: keptNodeIDs,
|
||||
removedNodeIDs: removedNodeIDs,
|
||||
portAdjusted: previousPort > 0 && previousPort != targetPort,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) executeTunnelDeleteForwardMigrations(plans []tunnelForwardMigrationPlan) ([]string, error, batchFailureDetail) {
|
||||
warnings := make([]string, 0)
|
||||
completed := make([]tunnelForwardMigrationPlan, 0, len(plans))
|
||||
|
||||
for _, plan := range plans {
|
||||
migrationWarnings, err := h.applyTunnelDeleteForwardMigration(plan)
|
||||
if err != nil {
|
||||
h.rollbackTunnelForwardMigrationPlans(completed)
|
||||
return warnings, err, batchFailureDetail{ID: plan.forward.ID, Name: plan.forward.Name, Reason: normalizeBatchFailureReason(errString(err))}
|
||||
}
|
||||
warnings = append(warnings, migrationWarnings...)
|
||||
completed = append(completed, plan)
|
||||
}
|
||||
|
||||
return warnings, nil, batchFailureDetail{}
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelDeleteForwardMigration(plan tunnelForwardMigrationPlan) ([]string, error) {
|
||||
if plan.forward == nil {
|
||||
return nil, errors.New("转发不存在")
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateForwardTunnel(plan.forward.ID, plan.targetTunnelID, time.Now().UnixMilli()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := h.replaceForwardPorts(plan.forward.ID, plan.targetTunnelID, plan.targetPort, ""); err != nil {
|
||||
h.rollbackForwardMutation(plan.forward, plan.oldPorts)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
updatedForward, err := h.getForwardRecord(plan.forward.ID)
|
||||
if err != nil {
|
||||
h.rollbackForwardMutation(plan.forward, plan.oldPorts)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
warnings := make([]string, 0)
|
||||
if len(plan.keptNodeIDs) > 0 {
|
||||
for _, nodeID := range plan.keptNodeIDs {
|
||||
if delErr := h.deleteForwardServicesOnNodeBatch(plan.forward, nodeID); delErr != nil {
|
||||
nodeLabel := fmt.Sprintf("%d", nodeID)
|
||||
if n, nErr := h.getNodeRecord(nodeID); nErr == nil && n != nil && strings.TrimSpace(n.Name) != "" {
|
||||
nodeLabel = strings.TrimSpace(n.Name)
|
||||
}
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %s 清理旧转发监听失败: %v", nodeLabel, delErr))
|
||||
}
|
||||
}
|
||||
time.Sleep(tunnelServiceBindRetryDelay)
|
||||
}
|
||||
|
||||
syncWarnings, err := h.syncForwardServicesWithWarnings(updatedForward, "UpdateService", true)
|
||||
if err != nil {
|
||||
h.rollbackForwardMutation(plan.forward, plan.oldPorts)
|
||||
return nil, err
|
||||
}
|
||||
warnings = append(warnings, syncWarnings...)
|
||||
|
||||
if len(plan.removedNodeIDs) > 0 {
|
||||
for _, nodeID := range plan.removedNodeIDs {
|
||||
if delErr := h.deleteForwardServicesOnNodeBatch(plan.forward, nodeID); delErr != nil {
|
||||
nodeLabel := fmt.Sprintf("%d", nodeID)
|
||||
if n, nErr := h.getNodeRecord(nodeID); nErr == nil && n != nil && strings.TrimSpace(n.Name) != "" {
|
||||
nodeLabel = strings.TrimSpace(n.Name)
|
||||
}
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %s 清理旧隧道残留服务失败: %v", nodeLabel, delErr))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return warnings, nil
|
||||
}
|
||||
|
||||
func (h *Handler) rollbackTunnelForwardMigrationPlans(plans []tunnelForwardMigrationPlan) {
|
||||
for i := len(plans) - 1; i >= 0; i-- {
|
||||
plan := plans[i]
|
||||
h.rollbackForwardMutation(plan.forward, plan.oldPorts)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) deleteTunnelAndCleanup(tunnelID int64) error {
|
||||
h.cleanupTunnelRuntime(tunnelID)
|
||||
h.cleanupFederationRuntime(tunnelID)
|
||||
if err := h.deleteTunnelByID(tunnelID); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func minInt(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestValidateTunnelEntryPortConflictsForNewEntriesDoesNotBlockOnSQLiteTx(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
})
|
||||
h := &Handler{repo: r}
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(name, secret, server_ip, port, created_time, status, tcp_listen_addr, udp_listen_addr, is_remote)
|
||||
VALUES
|
||||
('entry-old', 'secret-old', '10.0.0.1', '12000-12010', ?, 1, '[::]', '[::]', 0),
|
||||
('entry-new', 'secret-new', '10.0.0.2', '12000-12010', ?, 1, '[::]', '[::]', 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert nodes: %v", err)
|
||||
}
|
||||
var oldEntryID, newEntryID int64
|
||||
if err := r.DB().Raw(`SELECT id FROM node WHERE name = 'entry-old'`).Scan(&oldEntryID).Error; err != nil {
|
||||
t.Fatalf("load old entry id: %v", err)
|
||||
}
|
||||
if err := r.DB().Raw(`SELECT id FROM node WHERE name = 'entry-new'`).Scan(&newEntryID).Error; err != nil {
|
||||
t.Fatalf("load new entry id: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference)
|
||||
VALUES('sqlite-tunnel', 1, 1, 'tls', 1, ?, ?, 1, 1, '')
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
var tunnelID int64
|
||||
if err := r.DB().Raw(`SELECT id FROM tunnel WHERE name = 'sqlite-tunnel'`).Scan(&tunnelID).Error; err != nil {
|
||||
t.Fatalf("load tunnel id: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, inx, protocol)
|
||||
VALUES(?, '1', ?, 1, 'tls')
|
||||
`, tunnelID, oldEntryID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, created_time, updated_time, status, inx)
|
||||
VALUES(1, 'tester', 'forward-a', ?, '127.0.0.1:8080', 'fifo', ?, ?, 1, 1)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
var forwardID int64
|
||||
if err := r.DB().Raw(`SELECT id FROM forward WHERE name = 'forward-a'`).Scan(&forwardID).Error; err != nil {
|
||||
t.Fatalf("load forward id: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward_port(forward_id, node_id, port)
|
||||
VALUES(?, ?, 12001)
|
||||
`, forwardID, oldEntryID).Error; err != nil {
|
||||
t.Fatalf("insert forward_port: %v", err)
|
||||
}
|
||||
|
||||
tx := r.BeginTx()
|
||||
if tx == nil {
|
||||
t.Fatal("begin tx: nil transaction")
|
||||
}
|
||||
if tx.Error != nil {
|
||||
t.Fatalf("begin tx: %v", tx.Error)
|
||||
}
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
doneCh := make(chan struct{})
|
||||
go func() {
|
||||
defer close(doneCh)
|
||||
errCh <- h.validateTunnelEntryPortConflictsForNewEntriesTx(tx, tunnelID, []int64{oldEntryID}, []int64{oldEntryID, newEntryID})
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("unexpected validation error: %v", err)
|
||||
}
|
||||
case <-time.After(500 * time.Millisecond):
|
||||
_ = tx.Rollback().Error
|
||||
<-doneCh
|
||||
t.Fatal("validation blocked while transaction was open on sqlite")
|
||||
}
|
||||
|
||||
if err := tx.Rollback().Error; err != nil {
|
||||
t.Fatalf("rollback tx: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
type tunnelTrafficDelta struct {
|
||||
bytesIn int64
|
||||
bytesOut int64
|
||||
}
|
||||
|
||||
func unixMilliBucketMinute(nowMs int64) int64 {
|
||||
if nowMs <= 0 {
|
||||
return 0
|
||||
}
|
||||
const minuteMs = int64(time.Minute / time.Millisecond)
|
||||
return nowMs - (nowMs % minuteMs)
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowItem, nowMs int64) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if nodeID <= 0 || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
bucketTs := unixMilliBucketMinute(nowMs)
|
||||
if bucketTs <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
forwardDeltas := make(map[int64]tunnelTrafficDelta)
|
||||
var skippedParse, skippedZero int
|
||||
for _, item := range items {
|
||||
name := strings.TrimSpace(item.N)
|
||||
if name == "" || name == "web_api" {
|
||||
continue
|
||||
}
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(name)
|
||||
if !ok {
|
||||
skippedParse++
|
||||
continue
|
||||
}
|
||||
if item.D == 0 && item.U == 0 {
|
||||
skippedZero++
|
||||
continue
|
||||
}
|
||||
d := forwardDeltas[forwardID]
|
||||
d.bytesIn += item.D
|
||||
d.bytesOut += item.U
|
||||
forwardDeltas[forwardID] = d
|
||||
}
|
||||
if len(forwardDeltas) == 0 {
|
||||
if len(items) > 0 {
|
||||
log.Printf("monitoring debug op=tunnel_metric.no_forward_deltas node_id=%d items=%d skipped_parse=%d skipped_zero=%d", nodeID, len(items), skippedParse, skippedZero)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
forwardIDs := make([]int64, 0, len(forwardDeltas))
|
||||
for id := range forwardDeltas {
|
||||
forwardIDs = append(forwardIDs, id)
|
||||
}
|
||||
|
||||
forwardTunnelMap, err := h.repo.MapForwardIDsToTunnelIDs(forwardIDs)
|
||||
if err != nil {
|
||||
log.Printf("monitoring write skipped op=tunnel_metric.map_forward_to_tunnel node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
if len(forwardTunnelMap) == 0 {
|
||||
log.Printf("monitoring debug op=tunnel_metric.no_tunnel_map node_id=%d forward_ids=%v", nodeID, forwardIDs)
|
||||
return
|
||||
}
|
||||
|
||||
tunnelAgg := make(map[int64]tunnelTrafficDelta)
|
||||
for forwardID, delta := range forwardDeltas {
|
||||
tunnelID := forwardTunnelMap[forwardID]
|
||||
if tunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
a := tunnelAgg[tunnelID]
|
||||
a.bytesIn += delta.bytesIn
|
||||
a.bytesOut += delta.bytesOut
|
||||
tunnelAgg[tunnelID] = a
|
||||
}
|
||||
if len(tunnelAgg) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
|
||||
for tunnelID, delta := range tunnelAgg {
|
||||
if delta.bytesIn == 0 && delta.bytesOut == 0 {
|
||||
continue
|
||||
}
|
||||
metrics = append(metrics, &model.TunnelMetric{
|
||||
TunnelID: tunnelID,
|
||||
NodeID: nodeID,
|
||||
Timestamp: bucketTs,
|
||||
BytesIn: delta.bytesIn,
|
||||
BytesOut: delta.bytesOut,
|
||||
Connections: 0,
|
||||
Errors: 0,
|
||||
AvgLatencyMs: 0,
|
||||
})
|
||||
}
|
||||
if len(metrics) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
|
||||
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
|
||||
} else {
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,452 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
const (
|
||||
tunnelQualityProbeInterval = 1 * time.Second
|
||||
tunnelQualityProbeTimeout = 8 * time.Second
|
||||
tunnelQualityPingTimeoutMs = 5000
|
||||
tunnelQualityRetention = 24 * time.Hour // keep 24h of history
|
||||
tunnelQualityPruneInterval = 10 * time.Minute
|
||||
tunnelQualityReportInterval = 30 * time.Second // DB save interval
|
||||
)
|
||||
|
||||
type TunnelQualityHop struct {
|
||||
FromNodeID int64 `json:"fromNodeId"`
|
||||
FromNodeName string `json:"fromNodeName"`
|
||||
ToNodeID int64 `json:"toNodeId"`
|
||||
ToNodeName string `json:"toNodeName"`
|
||||
Latency float64 `json:"latency"`
|
||||
Loss float64 `json:"loss"`
|
||||
TargetIP string `json:"targetIp,omitempty"`
|
||||
TargetPort int `json:"targetPort,omitempty"`
|
||||
}
|
||||
|
||||
// tunnelQualitySnapshot is the in-memory latest probe result for a tunnel.
|
||||
type tunnelQualitySnapshot struct {
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
EntryToExitLatency float64 `json:"entryToExitLatency"`
|
||||
ExitToBingLatency float64 `json:"exitToBingLatency"`
|
||||
EntryToExitLoss float64 `json:"entryToExitLoss"`
|
||||
ExitToBingLoss float64 `json:"exitToBingLoss"`
|
||||
Success bool `json:"success"`
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
Timestamp int64 `json:"timestamp"`
|
||||
ChainDetails string `json:"chainDetails,omitempty"`
|
||||
|
||||
// internal fields for db reporting
|
||||
lastDBWrite int64 `json:"-"`
|
||||
}
|
||||
|
||||
// tunnelQualityProber runs periodic TCP ping probes against all enabled tunnels.
|
||||
// Design mirrors health.Checker: background goroutine with worker pool + scheduled cleanup.
|
||||
type tunnelQualityProber struct {
|
||||
handler *Handler
|
||||
cache sync.Map // tunnelID (int64) → *tunnelQualitySnapshot
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
interval time.Duration
|
||||
lastPrune int64
|
||||
probing int32 // atomic flag: 1 = probeAll running, 0 = idle
|
||||
}
|
||||
|
||||
// newTunnelQualityProber creates a new prober (not yet running).
|
||||
func newTunnelQualityProber(h *Handler) *tunnelQualityProber {
|
||||
return &tunnelQualityProber{
|
||||
handler: h,
|
||||
interval: tunnelQualityProbeInterval,
|
||||
}
|
||||
}
|
||||
|
||||
// Start launches the background probe loop (call from jobs.go).
|
||||
func (p *tunnelQualityProber) Start(ctx context.Context) {
|
||||
// Use the provided context so we stop with other background jobs.
|
||||
p.ctx, p.cancel = context.WithCancel(ctx)
|
||||
p.loop()
|
||||
}
|
||||
|
||||
// Stop halts the background probe loop.
|
||||
func (p *tunnelQualityProber) Stop() {
|
||||
if p == nil || p.cancel == nil {
|
||||
return
|
||||
}
|
||||
|
||||
p.cancel()
|
||||
}
|
||||
|
||||
// GetAll returns all cached quality snapshots (latest per tunnel).
|
||||
func (p *tunnelQualityProber) GetAll() []tunnelQualitySnapshot {
|
||||
var items []tunnelQualitySnapshot
|
||||
p.cache.Range(func(_, value interface{}) bool {
|
||||
if snap, ok := value.(*tunnelQualitySnapshot); ok {
|
||||
items = append(items, *snap)
|
||||
}
|
||||
return true
|
||||
})
|
||||
return items
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) loop() {
|
||||
// Initial delay to let the system boot up
|
||||
select {
|
||||
case <-time.After(5 * time.Second):
|
||||
case <-p.ctx.Done():
|
||||
return
|
||||
}
|
||||
|
||||
// Run once immediately
|
||||
p.probeAll()
|
||||
|
||||
ticker := time.NewTicker(p.interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-p.ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
p.probeAll()
|
||||
p.maybePrune()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) isEnabled() bool {
|
||||
if p == nil || p.handler == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
return p.handler.isTunnelQualityMonitoringEnabled()
|
||||
}
|
||||
|
||||
// maybePrune deletes old quality rows periodically (mirrors PruneServiceMonitorResults).
|
||||
func (p *tunnelQualityProber) maybePrune() {
|
||||
if !p.isEnabled() {
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if p.lastPrune > 0 && now-p.lastPrune < int64(tunnelQualityPruneInterval/time.Millisecond) {
|
||||
return
|
||||
}
|
||||
p.lastPrune = now
|
||||
|
||||
h := p.handler
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
cutoff := now - int64(tunnelQualityRetention/time.Millisecond)
|
||||
if err := h.repo.PruneTunnelQualityResults(cutoff); err != nil {
|
||||
log.Printf("tunnel_quality_prober: prune err=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeAll() {
|
||||
if !p.isEnabled() {
|
||||
return
|
||||
}
|
||||
|
||||
// Skip if previous probe round is still running (interval < timeout guard)
|
||||
if !atomic.CompareAndSwapInt32(&p.probing, 0, 1) {
|
||||
return
|
||||
}
|
||||
defer atomic.StoreInt32(&p.probing, 0)
|
||||
|
||||
h := p.handler
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
tunnelIDs, err := h.repo.ListEnabledTunnelIDs()
|
||||
if err != nil {
|
||||
log.Printf("tunnel_quality_prober: list enabled tunnels err=%v", err)
|
||||
return
|
||||
}
|
||||
if len(tunnelIDs) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Probe tunnels concurrently with a worker limit
|
||||
// (mirrors health.Checker worker pool pattern)
|
||||
const maxWorkers = 20
|
||||
sem := make(chan struct{}, maxWorkers)
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for _, tunnelID := range tunnelIDs {
|
||||
select {
|
||||
case <-p.ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
|
||||
wg.Add(1)
|
||||
sem <- struct{}{}
|
||||
go func(tid int64) {
|
||||
defer wg.Done()
|
||||
defer func() { <-sem }()
|
||||
p.probeTunnel(tid)
|
||||
}(tunnelID)
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
h := p.handler
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
snap := &tunnelQualitySnapshot{
|
||||
TunnelID: tunnelID,
|
||||
Timestamp: now,
|
||||
}
|
||||
|
||||
// Get tunnel chain info
|
||||
tunnel, err := h.getTunnelRecord(tunnelID)
|
||||
if err != nil {
|
||||
snap.ErrorMessage = "隧道不存在"
|
||||
p.storeResult(snap)
|
||||
return
|
||||
}
|
||||
|
||||
chainRows, err := h.listChainNodesForTunnel(tunnelID)
|
||||
if err != nil || len(chainRows) == 0 {
|
||||
snap.ErrorMessage = "隧道配置不完整"
|
||||
p.storeResult(snap)
|
||||
return
|
||||
}
|
||||
|
||||
ipPreference := h.repo.GetTunnelIPPreference(tunnelID)
|
||||
inNodes, midNodesGrouped, outNodes := splitChainNodeGroups(chainRows)
|
||||
|
||||
options := diagnosisExecOptions{
|
||||
commandTimeout: tunnelQualityProbeTimeout,
|
||||
pingTimeoutMS: tunnelQualityPingTimeoutMs,
|
||||
timeoutMessage: "探测超时",
|
||||
}
|
||||
|
||||
switch tunnel.Type {
|
||||
case 1:
|
||||
// Port forwarding: entry → Bing only
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(inNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
snap.Success = true
|
||||
} else {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
}
|
||||
case 2:
|
||||
// Tunnel forwarding: entry → exit + exit → Bing
|
||||
probeOK := true
|
||||
|
||||
if len(inNodes) > 0 && len(outNodes) > 0 {
|
||||
var hops []TunnelQualityHop
|
||||
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, outNodes[0])
|
||||
|
||||
for i := 0; i < len(nodesInPath)-1; i++ {
|
||||
source := nodesInPath[i]
|
||||
target := nodesInPath[i+1]
|
||||
|
||||
hop := TunnelQualityHop{
|
||||
FromNodeID: source.NodeID,
|
||||
FromNodeName: source.NodeName,
|
||||
ToNodeID: target.NodeID,
|
||||
ToNodeName: target.NodeName,
|
||||
}
|
||||
|
||||
targetNode, nodeErr := h.getNodeRecord(target.NodeID)
|
||||
if nodeErr != nil || targetNode == nil {
|
||||
snap.ErrorMessage = "节点 " + target.NodeName + " 不可用"
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
fromNode, _ := h.getNodeRecord(source.NodeID)
|
||||
targetIP, targetPort, resolveErr := resolveChainProbeTarget(fromNode, targetNode, target.Port, ipPreference, target.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
snap.ErrorMessage = "解析节点 " + target.NodeName + " 失败: " + resolveErr.Error()
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
hop.TargetIP = targetIP
|
||||
hop.TargetPort = targetPort
|
||||
|
||||
lat, loss, err := p.tcpPingNode(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)
|
||||
} else {
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
if snap.ErrorMessage == "" {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if probeOK {
|
||||
snap.EntryToExitLatency = totalLat
|
||||
snap.EntryToExitLoss = (1.0 - remainingSuccessProb) * 100.0
|
||||
} else {
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
}
|
||||
|
||||
if len(hops) > 0 {
|
||||
if b, err := json.Marshal(hops); err == nil {
|
||||
snap.ChainDetails = string(b)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Exit → Bing
|
||||
if len(outNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(outNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
} else {
|
||||
if snap.ErrorMessage == "" {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
probeOK = false
|
||||
}
|
||||
}
|
||||
|
||||
snap.Success = probeOK
|
||||
default:
|
||||
// Unknown type: entry → Bing
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(inNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
snap.Success = true
|
||||
} else {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
p.storeResult(snap)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) tcpPingNode(nodeID int64, ip string, port int, options diagnosisExecOptions) (latency float64, loss float64, err error) {
|
||||
h := p.handler
|
||||
if h == nil {
|
||||
return 0, 100, nil
|
||||
}
|
||||
|
||||
node, nodeErr := h.getNodeRecord(nodeID)
|
||||
if nodeErr != nil {
|
||||
return 0, 100, nodeErr
|
||||
}
|
||||
|
||||
var pingData map[string]interface{}
|
||||
var pingErr error
|
||||
if node != nil && node.IsRemote == 1 {
|
||||
pingData, pingErr = h.tcpPingViaRemoteNode(node, ip, port, options)
|
||||
} else {
|
||||
pingData, pingErr = h.tcpPingViaNode(nodeID, ip, port, options)
|
||||
}
|
||||
if pingErr != nil {
|
||||
return 0, 100, pingErr
|
||||
}
|
||||
|
||||
avgTime := asFloat(pingData["averageTime"], 0)
|
||||
packetLoss := asFloat(pingData["packetLoss"], 100)
|
||||
|
||||
return avgTime, packetLoss, nil
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) storeResult(snap *tunnelQualitySnapshot) {
|
||||
if snap == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Update in-memory cache (latest per tunnel)
|
||||
// Retain the lastDBWrite timestamp if it exists, so we only DB write every 30s
|
||||
var lastWrite int64
|
||||
if existing, ok := p.cache.Load(snap.TunnelID); ok {
|
||||
if eg, ok := existing.(*tunnelQualitySnapshot); ok {
|
||||
lastWrite = eg.lastDBWrite
|
||||
}
|
||||
}
|
||||
snap.lastDBWrite = lastWrite
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
writeToDB := false
|
||||
if now-snap.lastDBWrite >= int64(tunnelQualityReportInterval/time.Millisecond) {
|
||||
writeToDB = true
|
||||
snap.lastDBWrite = now
|
||||
}
|
||||
|
||||
p.cache.Store(snap.TunnelID, snap)
|
||||
|
||||
if !writeToDB {
|
||||
return
|
||||
}
|
||||
|
||||
// Persist to database (history)
|
||||
h := p.handler
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
successInt := 0
|
||||
if snap.Success {
|
||||
successInt = 1
|
||||
}
|
||||
|
||||
q := &model.TunnelQuality{
|
||||
TunnelID: snap.TunnelID,
|
||||
EntryToExitLatency: snap.EntryToExitLatency,
|
||||
ExitToBingLatency: snap.ExitToBingLatency,
|
||||
EntryToExitLoss: snap.EntryToExitLoss,
|
||||
ExitToBingLoss: snap.ExitToBingLoss,
|
||||
Success: successInt,
|
||||
ErrorMessage: snap.ErrorMessage,
|
||||
Timestamp: snap.Timestamp,
|
||||
ChainDetails: snap.ChainDetails,
|
||||
}
|
||||
if err := h.repo.InsertTunnelQuality(q); err != nil {
|
||||
log.Printf("tunnel_quality_prober: insert db err=%v tunnel_id=%d", err, snap.TunnelID)
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -19,8 +20,104 @@ const (
|
||||
githubHTMLBase = "https://github.com"
|
||||
upgradeTimeout = 5 * time.Minute
|
||||
batchWorkers = 5
|
||||
|
||||
releaseChannelStable = "stable"
|
||||
releaseChannelDev = "dev"
|
||||
)
|
||||
|
||||
var (
|
||||
stableVersionPattern = regexp.MustCompile(`^\d+(?:\.\d+)+$`)
|
||||
testKeywordPattern = regexp.MustCompile(`(?i)(alpha|beta|rc)`)
|
||||
)
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
}
|
||||
|
||||
func normalizeReleaseChannel(channel string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(channel)) {
|
||||
case releaseChannelDev:
|
||||
return releaseChannelDev
|
||||
default:
|
||||
return releaseChannelStable
|
||||
}
|
||||
}
|
||||
|
||||
func releaseChannelFromTag(tag string) string {
|
||||
normalized := strings.ToLower(strings.TrimSpace(tag))
|
||||
if normalized == "" {
|
||||
return releaseChannelDev
|
||||
}
|
||||
if testKeywordPattern.MatchString(normalized) {
|
||||
return releaseChannelDev
|
||||
}
|
||||
if stableVersionPattern.MatchString(normalized) {
|
||||
return releaseChannelStable
|
||||
}
|
||||
|
||||
return releaseChannelDev
|
||||
}
|
||||
|
||||
func releaseChannelLabel(channel string) string {
|
||||
if normalizeReleaseChannel(channel) == releaseChannelDev {
|
||||
return "测试版"
|
||||
}
|
||||
|
||||
return "正式版"
|
||||
}
|
||||
|
||||
func fetchGitHubReleases(perPage int) ([]githubRelease, error) {
|
||||
if perPage <= 0 {
|
||||
perPage = 20
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 15 * time.Second}
|
||||
resp, err := client.Get(fmt.Sprintf("%s/repos/%s/releases?per_page=%d", githubAPIBase, githubRepo, perPage))
|
||||
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 resolveLatestReleaseByChannel(channel string) (string, error) {
|
||||
normalizedChannel := normalizeReleaseChannel(channel)
|
||||
releases, err := fetchGitHubReleases(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 (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -30,6 +127,7 @@ func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
Version string `json:"version"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
@@ -40,12 +138,13 @@ func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestRelease()
|
||||
version, err = resolveLatestReleaseByChannel(channel)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新版本失败: %v", err)))
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新%s失败: %v", releaseChannelLabel(channel), err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -67,6 +166,7 @@ func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("升级失败: %v", err)))
|
||||
return
|
||||
}
|
||||
h.markNodePendingUpgradeRedeploy(req.ID)
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"version": version,
|
||||
@@ -75,61 +175,11 @@ func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func resolveLatestRelease() (string, error) {
|
||||
client := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
resp, err := client.Get(githubProxy + "/" + githubHTMLBase + "/" + githubRepo + "/releases/latest")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求GitHub失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusFound && resp.StatusCode != http.StatusMovedPermanently {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
location := resp.Header.Get("Location")
|
||||
if location == "" {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
parts := strings.Split(location, "/")
|
||||
tag := parts[len(parts)-1]
|
||||
if tag == "" || tag == "latest" {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
return tag, nil
|
||||
return resolveLatestReleaseByChannel(releaseChannelStable)
|
||||
}
|
||||
|
||||
func resolveLatestReleaseAPI() (string, error) {
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Get(githubAPIBase + "/repos/" + githubRepo + "/releases/latest")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求GitHub API失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
return "", fmt.Errorf("GitHub API返回 %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&release); err != nil {
|
||||
return "", fmt.Errorf("解析GitHub API响应失败: %v", err)
|
||||
}
|
||||
if strings.TrimSpace(release.TagName) == "" {
|
||||
return "", fmt.Errorf("无法从GitHub获取最新版本号")
|
||||
}
|
||||
|
||||
return release.TagName, nil
|
||||
return resolveLatestReleaseByChannel(releaseChannelStable)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -141,6 +191,7 @@ func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
var req struct {
|
||||
IDs []int64 `json:"ids"`
|
||||
Version string `json:"version"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
@@ -151,12 +202,13 @@ func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestRelease()
|
||||
version, err = resolveLatestReleaseByChannel(channel)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新版本失败: %v", err)))
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新%s失败: %v", releaseChannelLabel(channel), err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -195,6 +247,7 @@ func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
results[index] = upgradeResult{ID: nodeID, Success: false, Message: err.Error()}
|
||||
return
|
||||
}
|
||||
h.markNodePendingUpgradeRedeploy(nodeID)
|
||||
results[index] = upgradeResult{ID: nodeID, Success: true, Message: result.Message}
|
||||
}(i, id)
|
||||
}
|
||||
@@ -212,37 +265,28 @@ func (h *Handler) listReleases(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 15 * time.Second}
|
||||
resp, err := client.Get(githubAPIBase + "/repos/" + githubRepo + "/releases?per_page=20")
|
||||
var req struct {
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil && err != io.EOF {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
|
||||
releases, err := fetchGitHubReleases(50)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: GitHub API返回 %d: %s", resp.StatusCode, string(body))))
|
||||
return
|
||||
}
|
||||
|
||||
var releases []struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("解析版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
type releaseItem struct {
|
||||
Version string `json:"version"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"publishedAt"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
|
||||
items := make([]releaseItem, 0, len(releases))
|
||||
@@ -250,11 +294,20 @@ func (h *Handler) listReleases(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Draft {
|
||||
continue
|
||||
}
|
||||
tag := strings.TrimSpace(r.TagName)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
itemChannel := releaseChannelFromTag(tag)
|
||||
if itemChannel != channel {
|
||||
continue
|
||||
}
|
||||
items = append(items, releaseItem{
|
||||
Version: r.TagName,
|
||||
Version: tag,
|
||||
Name: r.Name,
|
||||
PublishedAt: r.PublishedAt,
|
||||
Prerelease: r.Prerelease,
|
||||
Prerelease: itemChannel == releaseChannelDev,
|
||||
Channel: itemChannel,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -289,3 +342,66 @@ func (h *Handler) nodeRollback(w http.ResponseWriter, r *http.Request) {
|
||||
"message": result.Message,
|
||||
}))
|
||||
}
|
||||
|
||||
func (h *Handler) markNodePendingUpgradeRedeploy(nodeID int64) {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return
|
||||
}
|
||||
h.upgradeMu.Lock()
|
||||
h.pendingUpgradeRedeploy[nodeID] = struct{}{}
|
||||
h.upgradeMu.Unlock()
|
||||
}
|
||||
|
||||
func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return false
|
||||
}
|
||||
h.upgradeMu.Lock()
|
||||
_, ok := h.pendingUpgradeRedeploy[nodeID]
|
||||
if ok {
|
||||
delete(h.pendingUpgradeRedeploy, nodeID)
|
||||
}
|
||||
h.upgradeMu.Unlock()
|
||||
return ok
|
||||
}
|
||||
|
||||
func (h *Handler) onNodeOnline(nodeID int64) {
|
||||
if !h.consumeNodePendingUpgradeRedeploy(nodeID) {
|
||||
return
|
||||
}
|
||||
h.redeployNodeRuntimeAfterUpgrade(nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
tunnelIDs, err := h.repo.ListActiveTunnelIDsByNode(nodeID)
|
||||
if err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: list tunnels for node %d failed: %v\n", nodeID, err)
|
||||
return
|
||||
}
|
||||
forwardIDs, err := h.repo.ListActiveForwardIDsByNode(nodeID)
|
||||
if err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: list forwards for node %d failed: %v\n", nodeID, err)
|
||||
return
|
||||
}
|
||||
|
||||
tunnelFailed := make(map[int64]struct{})
|
||||
for _, tunnelID := range tunnelIDs {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
tunnelFailed[tunnelID] = struct{}{}
|
||||
fmt.Printf("post-upgrade redeploy: tunnel %d failed on node %d: %v\n", tunnelID, nodeID, err)
|
||||
}
|
||||
}
|
||||
|
||||
for _, forwardID := range forwardIDs {
|
||||
forward, getErr := h.getForwardRecord(forwardID)
|
||||
if getErr != nil || forward == nil {
|
||||
continue
|
||||
}
|
||||
if _, skipped := tunnelFailed[forward.TunnelID]; skipped {
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: forward %d failed on node %d: %v\n", forwardID, nodeID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package handler
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestReleaseChannelFromTag(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
tag string
|
||||
expects string
|
||||
}{
|
||||
{name: "stable semantic version", tag: "2.1.4", expects: releaseChannelStable},
|
||||
{name: "v prefix should be dev", tag: "v2.1.4", expects: releaseChannelDev},
|
||||
{name: "rc release", tag: "2.1.4-rc2", expects: releaseChannelDev},
|
||||
{name: "beta release", tag: "2.1.4-beta.1", expects: releaseChannelDev},
|
||||
{name: "alpha release", tag: "2.1.4-alpha", expects: releaseChannelDev},
|
||||
{name: "non numeric tag", tag: "nightly", expects: releaseChannelDev},
|
||||
{name: "empty tag", tag: "", expects: releaseChannelDev},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := releaseChannelFromTag(tc.tag); got != tc.expects {
|
||||
t.Fatalf("releaseChannelFromTag(%q) = %q, want %q", tc.tag, got, tc.expects)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeReleaseChannel(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
expects string
|
||||
}{
|
||||
{input: "", expects: releaseChannelStable},
|
||||
{input: "stable", expects: releaseChannelStable},
|
||||
{input: "dev", expects: releaseChannelDev},
|
||||
{input: "DEV", expects: releaseChannelDev},
|
||||
{input: "preview", expects: releaseChannelStable},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
if got := normalizeReleaseChannel(tc.input); got != tc.expects {
|
||||
t.Fatalf("normalizeReleaseChannel(%q) = %q, want %q", tc.input, got, tc.expects)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func isUserQuotaExceeded(view *model.UserQuotaView) bool {
|
||||
if view == nil {
|
||||
return false
|
||||
}
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*bytesPerGB {
|
||||
return true
|
||||
}
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*bytesPerGB {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Handler) userQuotaBlockReason(userID int64, now int64) (string, error) {
|
||||
if h == nil || h.repo == nil || userID <= 0 {
|
||||
return "", nil
|
||||
}
|
||||
quota, err := h.repo.GetUserQuotaView(userID, time.UnixMilli(now))
|
||||
if err != nil || quota == nil {
|
||||
return "", err
|
||||
}
|
||||
if quota.DisabledByQuota == 1 || isUserQuotaExceeded(quota) {
|
||||
return "该用户流量配额已超额,禁止开启转发", nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
func (h *Handler) enforceUserQuotaIfNeeded(userID int64, quota *model.UserQuotaView) {
|
||||
if h == nil || h.repo == nil || userID <= 0 || quota == nil {
|
||||
return
|
||||
}
|
||||
if quota.DisabledByQuota == 1 || !isUserQuotaExceeded(quota) {
|
||||
return
|
||||
}
|
||||
|
||||
forwards, err := h.listActiveForwardsByUser(userID)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
pausedIDs := make([]int64, 0, len(forwards))
|
||||
now := time.Now().UnixMilli()
|
||||
for i := range forwards {
|
||||
forward := &forwards[i]
|
||||
if forward.Status != 1 {
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
|
||||
continue
|
||||
}
|
||||
if err := h.repo.UpdateForwardStatus(forward.ID, 0, now); err != nil {
|
||||
continue
|
||||
}
|
||||
pausedIDs = append(pausedIDs, forward.ID)
|
||||
}
|
||||
_ = h.repo.MarkUserQuotaDisabled(userID, pausedIDs, now)
|
||||
}
|
||||
|
||||
func (h *Handler) applyUserQuotaRelease(release *repo.UserQuotaRelease, now int64) {
|
||||
if h == nil || h.repo == nil || release == nil || release.UserID <= 0 || !release.UnblockUser {
|
||||
return
|
||||
}
|
||||
for _, forwardID := range release.ForwardIDs {
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil || forward == nil {
|
||||
continue
|
||||
}
|
||||
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.UpdateForwardStatus(forwardID, 1, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) resetUserQuotaWindows(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
releases, err := h.repo.RollUserQuotaWindows(now)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
for i := range releases {
|
||||
h.applyUserQuotaRelease(&releases[i], nowMs)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) userQuotaReset(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
var req struct {
|
||||
UserID int64 `json:"userId"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.UserID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("用户ID不能为空"))
|
||||
return
|
||||
}
|
||||
release, err := h.repo.ResetUserQuotaUsage(req.UserID, req.Scope, time.Now())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
nowMs := time.Now().UnixMilli()
|
||||
h.applyUserQuotaRelease(release, nowMs)
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) ensureUserForwardAllowedByQuota(userID int64, now int64) error {
|
||||
reason, err := h.userQuotaBlockReason(userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if reason != "" {
|
||||
return errors.New(reason)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -101,6 +101,10 @@ func shouldSkip(path string) bool {
|
||||
}
|
||||
|
||||
func requiresAdmin(path string) bool {
|
||||
if strings.HasPrefix(path, "/api/v1/monitor/permission/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/group/") {
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type SystemInfo struct {
|
||||
Uptime uint64 `json:"uptime"`
|
||||
BytesReceived uint64 `json:"bytes_received"`
|
||||
BytesTransmitted uint64 `json:"bytes_transmitted"`
|
||||
CPUUsage float64 `json:"cpu_usage"`
|
||||
MemoryUsage float64 `json:"memory_usage"`
|
||||
DiskUsage float64 `json:"disk_usage"`
|
||||
Load1 float64 `json:"load1"`
|
||||
Load5 float64 `json:"load5"`
|
||||
Load15 float64 `json:"load15"`
|
||||
TCPConns int64 `json:"tcp_conns"`
|
||||
UDPConns int64 `json:"udp_conns"`
|
||||
NetInSpeed int64 `json:"net_in_speed"`
|
||||
NetOutSpeed int64 `json:"net_out_speed"`
|
||||
}
|
||||
|
||||
type IngestionService struct {
|
||||
repo *repo.Repository
|
||||
nodeBuffer []*model.NodeMetric
|
||||
nodeBufferMu sync.Mutex
|
||||
flushInterval time.Duration
|
||||
retentionDays int
|
||||
}
|
||||
|
||||
func NewIngestionService(repo *repo.Repository) *IngestionService {
|
||||
return &IngestionService{
|
||||
repo: repo,
|
||||
nodeBuffer: make([]*model.NodeMetric, 0, 500),
|
||||
flushInterval: 30 * time.Second,
|
||||
retentionDays: 7,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *IngestionService) Start(ctx context.Context) {
|
||||
flushTicker := time.NewTicker(s.flushInterval)
|
||||
defer flushTicker.Stop()
|
||||
|
||||
pruneTicker := time.NewTicker(1 * time.Hour)
|
||||
defer pruneTicker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
s.flushNodeMetrics()
|
||||
return
|
||||
case <-flushTicker.C:
|
||||
s.flushNodeMetrics()
|
||||
case <-pruneTicker.C:
|
||||
s.pruneMetrics()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *IngestionService) RecordNodeMetric(nodeID int64, info SystemInfo) {
|
||||
m := &model.NodeMetric{
|
||||
NodeID: nodeID,
|
||||
Timestamp: time.Now().UnixMilli(),
|
||||
CPUUsage: info.CPUUsage,
|
||||
MemUsage: info.MemoryUsage,
|
||||
DiskUsage: info.DiskUsage,
|
||||
NetInBytes: int64(info.BytesReceived),
|
||||
NetOutBytes: int64(info.BytesTransmitted),
|
||||
NetInSpeed: info.NetInSpeed,
|
||||
NetOutSpeed: info.NetOutSpeed,
|
||||
Load1: info.Load1,
|
||||
Load5: info.Load5,
|
||||
Load15: info.Load15,
|
||||
TCPConns: info.TCPConns,
|
||||
UDPConns: info.UDPConns,
|
||||
Uptime: int64(info.Uptime),
|
||||
}
|
||||
|
||||
s.nodeBufferMu.Lock()
|
||||
s.nodeBuffer = append(s.nodeBuffer, m)
|
||||
shouldFlush := len(s.nodeBuffer) >= 200
|
||||
s.nodeBufferMu.Unlock()
|
||||
|
||||
if shouldFlush {
|
||||
go s.flushNodeMetrics()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *IngestionService) flushNodeMetrics() {
|
||||
s.nodeBufferMu.Lock()
|
||||
if len(s.nodeBuffer) == 0 {
|
||||
s.nodeBufferMu.Unlock()
|
||||
return
|
||||
}
|
||||
buffer := s.nodeBuffer
|
||||
s.nodeBuffer = make([]*model.NodeMetric, 0, 500)
|
||||
s.nodeBufferMu.Unlock()
|
||||
|
||||
if s.repo == nil {
|
||||
return
|
||||
}
|
||||
if err := s.repo.InsertNodeMetricBatch(buffer); err != nil {
|
||||
log.Printf("monitoring write failed op=node_metric.flush count=%d err=%v", len(buffer), err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetrics() {
|
||||
cutoff := time.Now().Add(-time.Duration(s.retentionDays) * 24 * time.Hour).UnixMilli()
|
||||
if s.repo == nil {
|
||||
return
|
||||
}
|
||||
if err := s.repo.PruneNodeMetrics(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=node_metric cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
if err := s.repo.PruneTunnelMetrics(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=tunnel_metric cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
if err := s.repo.PruneServiceMonitorResults(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=service_monitor_result cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *IngestionService) GetLatestMetric(nodeID int64) (*model.NodeMetric, error) {
|
||||
return s.repo.GetLatestNodeMetric(nodeID)
|
||||
}
|
||||
|
||||
func (s *IngestionService) GetMetrics(nodeID int64, startMs, endMs int64) ([]model.NodeMetric, error) {
|
||||
return s.repo.GetNodeMetrics(nodeID, startMs, endMs)
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestRecordNodeMetric(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
|
||||
info := SystemInfo{
|
||||
Uptime: 86400,
|
||||
BytesReceived: 1024000,
|
||||
BytesTransmitted: 2048000,
|
||||
CPUUsage: 45.5,
|
||||
MemoryUsage: 60.2,
|
||||
DiskUsage: 30.1,
|
||||
Load1: 1.5,
|
||||
Load5: 1.2,
|
||||
Load15: 0.9,
|
||||
TCPConns: 100,
|
||||
UDPConns: 50,
|
||||
NetInSpeed: 51200,
|
||||
NetOutSpeed: 102400,
|
||||
}
|
||||
|
||||
svc.RecordNodeMetric(1, info)
|
||||
svc.flushNodeMetrics()
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, time.Now().UnixMilli()-60000, time.Now().UnixMilli()+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 {
|
||||
t.Fatalf("expected 1 metric, got %d", len(metrics))
|
||||
}
|
||||
|
||||
m := metrics[0]
|
||||
if m.CPUUsage != 45.5 {
|
||||
t.Fatalf("expected CPUUsage 45.5, got %f", m.CPUUsage)
|
||||
}
|
||||
if m.MemUsage != 60.2 {
|
||||
t.Fatalf("expected MemUsage 60.2, got %f", m.MemUsage)
|
||||
}
|
||||
if m.DiskUsage != 30.1 {
|
||||
t.Fatalf("expected DiskUsage 30.1, got %f", m.DiskUsage)
|
||||
}
|
||||
if m.Load1 != 1.5 {
|
||||
t.Fatalf("expected Load1 1.5, got %f", m.Load1)
|
||||
}
|
||||
if m.TCPConns != 100 {
|
||||
t.Fatalf("expected TCPConns 100, got %d", m.TCPConns)
|
||||
}
|
||||
if m.UDPConns != 50 {
|
||||
t.Fatalf("expected UDPConns 50, got %d", m.UDPConns)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordNodeMetricAutoFlush(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
|
||||
info := SystemInfo{
|
||||
CPUUsage: 50.0,
|
||||
MemoryUsage: 60.0,
|
||||
DiskUsage: 30.0,
|
||||
}
|
||||
|
||||
for i := 0; i < 250; i++ {
|
||||
svc.RecordNodeMetric(1, info)
|
||||
}
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, time.Now().UnixMilli()-60000, time.Now().UnixMilli()+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics: %v", err)
|
||||
}
|
||||
if len(metrics) < 200 {
|
||||
t.Fatalf("expected at least 200 metrics after auto-flush, got %d", len(metrics))
|
||||
}
|
||||
}
|
||||
|
||||
func TestIngestionServiceStart(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.flushInterval = 100 * time.Millisecond
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
info := SystemInfo{
|
||||
CPUUsage: 45.0,
|
||||
MemoryUsage: 55.0,
|
||||
DiskUsage: 35.0,
|
||||
}
|
||||
|
||||
go svc.Start(ctx)
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
svc.RecordNodeMetric(1, info)
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
|
||||
<-ctx.Done()
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, time.Now().UnixMilli()-60000, time.Now().UnixMilli()+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics: %v", err)
|
||||
}
|
||||
if len(metrics) == 0 {
|
||||
t.Fatalf("expected metrics after service run")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLatestMetric(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
info1 := SystemInfo{CPUUsage: 40.0, MemoryUsage: 50.0, DiskUsage: 30.0}
|
||||
svc.RecordNodeMetric(1, info1)
|
||||
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
|
||||
info2 := SystemInfo{CPUUsage: 60.0, MemoryUsage: 70.0, DiskUsage: 40.0}
|
||||
svc.RecordNodeMetric(1, info2)
|
||||
|
||||
svc.flushNodeMetrics()
|
||||
|
||||
latest, err := svc.GetLatestMetric(1)
|
||||
if err != nil {
|
||||
t.Fatalf("get latest: %v", err)
|
||||
}
|
||||
if latest == nil {
|
||||
t.Fatalf("expected latest metric")
|
||||
}
|
||||
if latest.CPUUsage != 60.0 {
|
||||
t.Fatalf("expected latest CPUUsage 60.0, got %f", latest.CPUUsage)
|
||||
}
|
||||
|
||||
_ = now
|
||||
|
||||
latestNone, err := svc.GetLatestMetric(999)
|
||||
if err != nil {
|
||||
t.Fatalf("get latest for non-existent: %v", err)
|
||||
}
|
||||
if latestNone != nil {
|
||||
t.Fatalf("expected nil for non-existent node")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMetricsWithTimeRange(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
for i := 0; i < 5; i++ {
|
||||
info := SystemInfo{
|
||||
CPUUsage: float64(40 + i*5),
|
||||
MemoryUsage: 50.0,
|
||||
DiskUsage: 30.0,
|
||||
}
|
||||
svc.RecordNodeMetric(1, info)
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
|
||||
svc.flushNodeMetrics()
|
||||
|
||||
metrics, err := svc.GetMetrics(1, now-60000, now+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 5 {
|
||||
t.Fatalf("expected 5 metrics, got %d", len(metrics))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneMetrics(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.retentionDays = 1
|
||||
|
||||
info := SystemInfo{CPUUsage: 50.0, MemoryUsage: 60.0, DiskUsage: 30.0}
|
||||
|
||||
svc.RecordNodeMetric(1, info)
|
||||
svc.flushNodeMetrics()
|
||||
|
||||
svc.pruneMetrics()
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, time.Now().UnixMilli()-60000, time.Now().UnixMilli()+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 {
|
||||
t.Fatalf("expected 1 metric (not pruned), got %d", len(metrics))
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultipleNodes(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
|
||||
info := SystemInfo{
|
||||
CPUUsage: 50.0,
|
||||
MemoryUsage: 60.0,
|
||||
DiskUsage: 30.0,
|
||||
}
|
||||
|
||||
svc.RecordNodeMetric(1, info)
|
||||
svc.RecordNodeMetric(2, info)
|
||||
svc.RecordNodeMetric(3, info)
|
||||
|
||||
svc.flushNodeMetrics()
|
||||
|
||||
for nodeID := int64(1); nodeID <= 3; nodeID++ {
|
||||
metrics, err := r.GetNodeMetrics(nodeID, time.Now().UnixMilli()-60000, time.Now().UnixMilli()+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics for node %d: %v", nodeID, err)
|
||||
}
|
||||
if len(metrics) != 1 {
|
||||
t.Fatalf("expected 1 metric for node %d, got %d", nodeID, len(metrics))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestZeroValues(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
|
||||
info := SystemInfo{}
|
||||
|
||||
svc.RecordNodeMetric(1, info)
|
||||
svc.flushNodeMetrics()
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, time.Now().UnixMilli()-60000, time.Now().UnixMilli()+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 {
|
||||
t.Fatalf("expected 1 metric, got %d", len(metrics))
|
||||
}
|
||||
|
||||
m := metrics[0]
|
||||
if m.CPUUsage != 0 || m.MemUsage != 0 || m.DiskUsage != 0 {
|
||||
t.Fatalf("expected zero values, got CPU=%f Mem=%f Disk=%f", m.CPUUsage, m.MemUsage, m.DiskUsage)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type ServiceMonitorLimits struct {
|
||||
CheckerScanIntervalSec int `json:"checkerScanIntervalSec"`
|
||||
WorkerLimit int `json:"workerLimit"`
|
||||
|
||||
MinIntervalSec int `json:"minIntervalSec"`
|
||||
DefaultIntervalSec int `json:"defaultIntervalSec"`
|
||||
|
||||
MinTimeoutSec int `json:"minTimeoutSec"`
|
||||
DefaultTimeoutSec int `json:"defaultTimeoutSec"`
|
||||
MaxTimeoutSec int `json:"maxTimeoutSec"`
|
||||
}
|
||||
|
||||
const (
|
||||
ConfigServiceMonitorCheckerScanIntervalSec = "service_monitor_checker_scan_interval_sec"
|
||||
ConfigServiceMonitorWorkerLimit = "service_monitor_worker_limit"
|
||||
ConfigServiceMonitorMinIntervalSec = "service_monitor_min_interval_sec"
|
||||
ConfigServiceMonitorDefaultIntervalSec = "service_monitor_default_interval_sec"
|
||||
ConfigServiceMonitorMinTimeoutSec = "service_monitor_min_timeout_sec"
|
||||
ConfigServiceMonitorDefaultTimeoutSec = "service_monitor_default_timeout_sec"
|
||||
ConfigServiceMonitorMaxTimeoutSec = "service_monitor_max_timeout_sec"
|
||||
)
|
||||
|
||||
func DefaultServiceMonitorLimits() ServiceMonitorLimits {
|
||||
return ServiceMonitorLimits{
|
||||
CheckerScanIntervalSec: 1,
|
||||
WorkerLimit: 20,
|
||||
MinIntervalSec: 1,
|
||||
DefaultIntervalSec: 1,
|
||||
MinTimeoutSec: 1,
|
||||
DefaultTimeoutSec: 5,
|
||||
MaxTimeoutSec: 60,
|
||||
}
|
||||
}
|
||||
|
||||
// ServiceMonitorLimitsFromConfigMap parses limits from vite_config values.
|
||||
// Missing/invalid values fall back to defaults.
|
||||
func ServiceMonitorLimitsFromConfigMap(cfg map[string]string) ServiceMonitorLimits {
|
||||
limits := DefaultServiceMonitorLimits()
|
||||
if cfg == nil {
|
||||
return limits
|
||||
}
|
||||
|
||||
limits.CheckerScanIntervalSec = parseConfigInt(cfg, ConfigServiceMonitorCheckerScanIntervalSec, limits.CheckerScanIntervalSec)
|
||||
limits.WorkerLimit = parseConfigInt(cfg, ConfigServiceMonitorWorkerLimit, limits.WorkerLimit)
|
||||
limits.MinIntervalSec = parseConfigInt(cfg, ConfigServiceMonitorMinIntervalSec, limits.MinIntervalSec)
|
||||
limits.DefaultIntervalSec = parseConfigInt(cfg, ConfigServiceMonitorDefaultIntervalSec, limits.DefaultIntervalSec)
|
||||
limits.MinTimeoutSec = parseConfigInt(cfg, ConfigServiceMonitorMinTimeoutSec, limits.MinTimeoutSec)
|
||||
limits.DefaultTimeoutSec = parseConfigInt(cfg, ConfigServiceMonitorDefaultTimeoutSec, limits.DefaultTimeoutSec)
|
||||
limits.MaxTimeoutSec = parseConfigInt(cfg, ConfigServiceMonitorMaxTimeoutSec, limits.MaxTimeoutSec)
|
||||
|
||||
return normalizeServiceMonitorLimits(limits)
|
||||
}
|
||||
|
||||
func normalizeServiceMonitorLimits(limits ServiceMonitorLimits) ServiceMonitorLimits {
|
||||
if limits.CheckerScanIntervalSec <= 0 {
|
||||
limits.CheckerScanIntervalSec = 30
|
||||
}
|
||||
if limits.WorkerLimit <= 0 {
|
||||
limits.WorkerLimit = 5
|
||||
}
|
||||
if limits.WorkerLimit > 50 {
|
||||
limits.WorkerLimit = 50
|
||||
}
|
||||
|
||||
if limits.MinIntervalSec <= 0 {
|
||||
limits.MinIntervalSec = limits.CheckerScanIntervalSec
|
||||
}
|
||||
if limits.MinIntervalSec < limits.CheckerScanIntervalSec {
|
||||
limits.MinIntervalSec = limits.CheckerScanIntervalSec
|
||||
}
|
||||
if limits.DefaultIntervalSec <= 0 {
|
||||
limits.DefaultIntervalSec = 60
|
||||
}
|
||||
if limits.DefaultIntervalSec < limits.MinIntervalSec {
|
||||
limits.DefaultIntervalSec = limits.MinIntervalSec
|
||||
}
|
||||
|
||||
if limits.MinTimeoutSec <= 0 {
|
||||
limits.MinTimeoutSec = 1
|
||||
}
|
||||
if limits.DefaultTimeoutSec <= 0 {
|
||||
limits.DefaultTimeoutSec = 5
|
||||
}
|
||||
if limits.DefaultTimeoutSec < limits.MinTimeoutSec {
|
||||
limits.DefaultTimeoutSec = limits.MinTimeoutSec
|
||||
}
|
||||
if limits.MaxTimeoutSec <= 0 {
|
||||
limits.MaxTimeoutSec = 60
|
||||
}
|
||||
if limits.MaxTimeoutSec < limits.DefaultTimeoutSec {
|
||||
limits.MaxTimeoutSec = limits.DefaultTimeoutSec
|
||||
}
|
||||
|
||||
return limits
|
||||
}
|
||||
|
||||
func parseConfigInt(cfg map[string]string, key string, fallback int) int {
|
||||
v := strings.TrimSpace(cfg[key])
|
||||
if v == "" {
|
||||
return fallback
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return fallback
|
||||
}
|
||||
return n
|
||||
}
|
||||
@@ -29,68 +29,75 @@ 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:"column:in_flow;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"`
|
||||
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"`
|
||||
}
|
||||
|
||||
func (Forward) TableName() string { return "forward" }
|
||||
|
||||
type ForwardPort struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null"`
|
||||
NodeID int64 `gorm:"column:node_id;not null"`
|
||||
Port int `gorm:"not null"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null"`
|
||||
NodeID int64 `gorm:"column:node_id;not null"`
|
||||
Port int `gorm:"not null"`
|
||||
InIP sql.NullString `gorm:"column:in_ip;type:text"`
|
||||
}
|
||||
|
||||
func (ForwardPort) TableName() string { return "forward_port" }
|
||||
|
||||
type Node struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
Secret string `gorm:"type:varchar(100);not null"`
|
||||
ServerIP string `gorm:"column:server_ip;type:varchar(100);not null"`
|
||||
ServerIPV4 sql.NullString `gorm:"column:server_ip_v4;type:varchar(100)"`
|
||||
ServerIPV6 sql.NullString `gorm:"column:server_ip_v6;type:varchar(100)"`
|
||||
Port string `gorm:"type:text;not null"`
|
||||
InterfaceName sql.NullString `gorm:"column:interface_name;type:varchar(200)"`
|
||||
Version sql.NullString `gorm:"type:varchar(100)"`
|
||||
HTTP int `gorm:"column:http;not null;default:0"`
|
||||
TLS int `gorm:"column:tls;not null;default:0"`
|
||||
Socks int `gorm:"not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
TCPListenAddr string `gorm:"column:tcp_listen_addr;type:varchar(100);not null;default:'[::]'"`
|
||||
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"`
|
||||
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"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
Remark sql.NullString `gorm:"column:remark;type:text"`
|
||||
ExpiryTime sql.NullInt64 `gorm:"column:expiry_time"`
|
||||
RenewalCycle sql.NullString `gorm:"column:renewal_cycle;type:varchar(20)"`
|
||||
Secret string `gorm:"type:varchar(100);not null"`
|
||||
ServerIP string `gorm:"column:server_ip;type:varchar(100);not null"`
|
||||
ServerIPV4 sql.NullString `gorm:"column:server_ip_v4;type:varchar(100)"`
|
||||
ServerIPV6 sql.NullString `gorm:"column:server_ip_v6;type:varchar(100)"`
|
||||
ExtraIPs sql.NullString `gorm:"column:extra_ips;type:text"`
|
||||
Port string `gorm:"type:text;not null"`
|
||||
InterfaceName sql.NullString `gorm:"column:interface_name;type:varchar(200)"`
|
||||
Version sql.NullString `gorm:"type:varchar(100)"`
|
||||
HTTP int `gorm:"column:http;not null;default:0"`
|
||||
TLS int `gorm:"column:tls;not null;default:0"`
|
||||
Socks int `gorm:"not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
TCPListenAddr string `gorm:"column:tcp_listen_addr;type:varchar(100);not null;default:'[::]'"`
|
||||
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"`
|
||||
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"`
|
||||
ExpiryReminderDismissed int `gorm:"column:expiry_reminder_dismissed;not null;default:0"`
|
||||
}
|
||||
|
||||
func (Node) TableName() string { return "node" }
|
||||
|
||||
type SpeedLimit struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
Speed int `gorm:"not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
TunnelName string `gorm:"column:tunnel_name;type:varchar(100);not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
Speed int `gorm:"not null"`
|
||||
TunnelID sql.NullInt64 `gorm:"column:tunnel_id"`
|
||||
TunnelName sql.NullString `gorm:"column:tunnel_name;type:varchar(100)"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
|
||||
Status int `gorm:"not null"`
|
||||
}
|
||||
|
||||
func (SpeedLimit) TableName() string { return "speed_limit" }
|
||||
@@ -123,6 +130,23 @@ type Tunnel struct {
|
||||
|
||||
func (Tunnel) TableName() string { return "tunnel" }
|
||||
|
||||
type UserQuota struct {
|
||||
UserID int64 `gorm:"column:user_id;primaryKey"`
|
||||
DailyLimitGB int64 `gorm:"column:daily_limit_gb;not null;default:0"`
|
||||
MonthlyLimitGB int64 `gorm:"column:monthly_limit_gb;not null;default:0"`
|
||||
DailyUsedBytes int64 `gorm:"column:daily_used_bytes;not null;default:0"`
|
||||
MonthlyUsedBytes int64 `gorm:"column:monthly_used_bytes;not null;default:0"`
|
||||
DayKey int64 `gorm:"column:day_key;not null;default:0"`
|
||||
MonthKey int64 `gorm:"column:month_key;not null;default:0"`
|
||||
DisabledByQuota int `gorm:"column:disabled_by_quota;not null;default:0"`
|
||||
DisabledAt int64 `gorm:"column:disabled_at;not null;default:0"`
|
||||
PausedForwardIDs string `gorm:"column:paused_forward_ids;type:text;not null;default:''"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (UserQuota) TableName() string { return "user_quota" }
|
||||
|
||||
type ChainTunnel struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
@@ -132,6 +156,7 @@ type ChainTunnel struct {
|
||||
Strategy sql.NullString `gorm:"type:varchar(10)"`
|
||||
Inx sql.NullInt64 `gorm:"column:inx"`
|
||||
Protocol sql.NullString `gorm:"type:varchar(10)"`
|
||||
ConnectIP sql.NullString `gorm:"column:connect_ip;type:varchar(45)"`
|
||||
}
|
||||
|
||||
func (ChainTunnel) TableName() string { return "chain_tunnel" }
|
||||
@@ -210,10 +235,20 @@ type GroupPermissionGrant struct {
|
||||
|
||||
func (GroupPermissionGrant) TableName() string { return "group_permission_grant" }
|
||||
|
||||
// MonitorPermission grants a non-admin user access to monitoring endpoints.
|
||||
// One row per user_id.
|
||||
type MonitorPermission struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
UserID int64 `gorm:"column:user_id;not null;uniqueIndex:idx_monitor_permission_user" json:"userId"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null" json:"createdTime"`
|
||||
}
|
||||
|
||||
func (MonitorPermission) TableName() string { return "monitor_permission" }
|
||||
|
||||
type ViteConfig struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
Name string `gorm:"type:varchar(200);not null;uniqueIndex" json:"name"`
|
||||
Value string `gorm:"type:varchar(200);not null" json:"value"`
|
||||
Value string `gorm:"type:text;not null" json:"value"`
|
||||
Time int64 `gorm:"not null" json:"time"`
|
||||
}
|
||||
|
||||
@@ -314,28 +349,36 @@ type BackupData struct {
|
||||
}
|
||||
|
||||
type UserBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
User string `json:"user"`
|
||||
Pwd string `json:"pwd"`
|
||||
RoleID int `json:"roleId"`
|
||||
ExpTime int64 `json:"expTime"`
|
||||
Flow int64 `json:"flow"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
Num int `json:"num"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
||||
Status int `json:"status"`
|
||||
ID int64 `json:"id"`
|
||||
User string `json:"user"`
|
||||
Pwd string `json:"pwd"`
|
||||
RoleID int `json:"roleId"`
|
||||
ExpTime int64 `json:"expTime"`
|
||||
Flow int64 `json:"flow"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
DailyQuotaGB int64 `json:"dailyQuotaGB,omitempty"`
|
||||
MonthlyQuotaGB int64 `json:"monthlyQuotaGB,omitempty"`
|
||||
DisabledByQuota int `json:"disabledByQuota,omitempty"`
|
||||
QuotaDisabledAt int64 `json:"quotaDisabledAt,omitempty"`
|
||||
Num int `json:"num"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
||||
Status int `json:"status"`
|
||||
}
|
||||
|
||||
type NodeBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Remark string `json:"remark,omitempty"`
|
||||
ExpiryTime int64 `json:"expiryTime,omitempty"`
|
||||
RenewalCycle string `json:"renewalCycle,omitempty"`
|
||||
Secret string `json:"secret"`
|
||||
ServerIP string `json:"serverIp"`
|
||||
ServerIPv4 string `json:"serverIpV4,omitempty"`
|
||||
ServerIPv6 string `json:"serverIpV6,omitempty"`
|
||||
ExtraIPs string `json:"extraIPs,omitempty"`
|
||||
Port string `json:"port"`
|
||||
InterfaceName string `json:"interfaceName,omitempty"`
|
||||
Version string `json:"version,omitempty"`
|
||||
@@ -395,6 +438,7 @@ type ForwardBackup struct {
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
Inx int `json:"inx"`
|
||||
SpeedID *int64 `json:"speedId,omitempty"`
|
||||
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||
}
|
||||
|
||||
@@ -421,8 +465,8 @@ type SpeedLimitBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Speed int64 `json:"speed"`
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
TunnelName string `json:"tunnelName"`
|
||||
TunnelID *int64 `json:"tunnelId,omitempty"`
|
||||
TunnelName string `json:"tunnelName,omitempty"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime,omitempty"`
|
||||
Status int `json:"status"`
|
||||
@@ -492,6 +536,7 @@ type ForwardRecord struct {
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
}
|
||||
|
||||
// TunnelRecord is a minimal tunnel view used by control plane.
|
||||
@@ -503,10 +548,24 @@ type TunnelRecord struct {
|
||||
TrafficRatio float64
|
||||
}
|
||||
|
||||
type UserQuotaView struct {
|
||||
UserID int64
|
||||
DailyLimitGB int64
|
||||
MonthlyLimitGB int64
|
||||
DailyUsedBytes int64
|
||||
MonthlyUsedBytes int64
|
||||
DayKey int64
|
||||
MonthKey int64
|
||||
DisabledByQuota int
|
||||
DisabledAt int64
|
||||
PausedForwardIDs string
|
||||
}
|
||||
|
||||
// ForwardPortRecord is a forward port mapping used by control plane.
|
||||
type ForwardPortRecord struct {
|
||||
NodeID int64
|
||||
Port int
|
||||
InIP string
|
||||
}
|
||||
|
||||
// NodeRecord is a node view used by control plane.
|
||||
@@ -516,6 +575,7 @@ type NodeRecord struct {
|
||||
ServerIP string
|
||||
ServerIPv4 string
|
||||
ServerIPv6 string
|
||||
ExtraIPs string
|
||||
Status int
|
||||
PortRange string
|
||||
TCPListenAddr string
|
||||
@@ -535,6 +595,7 @@ type ChainNodeRecord struct {
|
||||
NodeName string
|
||||
Protocol string
|
||||
Strategy string
|
||||
ConnectIP string
|
||||
}
|
||||
|
||||
type UserTunnelLimiterInfo struct {
|
||||
@@ -563,6 +624,7 @@ type UserTunnelDetail struct {
|
||||
UserID int64
|
||||
TunnelID int64
|
||||
TunnelName string
|
||||
Status int
|
||||
TunnelFlow int
|
||||
Flow int64
|
||||
InFlow int64
|
||||
@@ -589,3 +651,84 @@ type UserForwardDetail struct {
|
||||
Status int
|
||||
CreatedAt int64
|
||||
}
|
||||
|
||||
type NodeMetric struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;index:idx_node_metric_node_time,priority:1" json:"nodeId"`
|
||||
Timestamp int64 `gorm:"not null;index:idx_node_metric_node_time,priority:2;index:idx_node_metric_time" json:"timestamp"`
|
||||
CPUUsage float64 `gorm:"column:cpu_usage" json:"cpuUsage"`
|
||||
MemUsage float64 `gorm:"column:mem_usage" json:"memoryUsage"`
|
||||
DiskUsage float64 `gorm:"column:disk_usage" json:"diskUsage"`
|
||||
NetInBytes int64 `gorm:"column:net_in_bytes" json:"netInBytes"`
|
||||
NetOutBytes int64 `gorm:"column:net_out_bytes" json:"netOutBytes"`
|
||||
NetInSpeed int64 `gorm:"column:net_in_speed" json:"netInSpeed"`
|
||||
NetOutSpeed int64 `gorm:"column:net_out_speed" json:"netOutSpeed"`
|
||||
Load1 float64 `gorm:"column:load1" json:"load1"`
|
||||
Load5 float64 `gorm:"column:load5" json:"load5"`
|
||||
Load15 float64 `gorm:"column:load15" json:"load15"`
|
||||
TCPConns int64 `gorm:"column:tcp_conns" json:"tcpConns"`
|
||||
UDPConns int64 `gorm:"column:udp_conns" json:"udpConns"`
|
||||
Uptime int64 `gorm:"column:uptime" json:"uptime"`
|
||||
}
|
||||
|
||||
func (NodeMetric) TableName() string { return "node_metric" }
|
||||
|
||||
type TunnelMetric struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null;uniqueIndex:idx_tunnel_metric_tunnel_time,priority:1" json:"tunnelId"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_tunnel_metric_tunnel_time,priority:2" json:"nodeId"`
|
||||
Timestamp int64 `gorm:"not null;uniqueIndex:idx_tunnel_metric_tunnel_time,priority:3;index:idx_tunnel_metric_time" json:"timestamp"`
|
||||
BytesIn int64 `gorm:"column:bytes_in" json:"bytesIn"`
|
||||
BytesOut int64 `gorm:"column:bytes_out" json:"bytesOut"`
|
||||
Connections int64 `gorm:"column:connections" json:"connections"`
|
||||
Errors int64 `gorm:"column:errors" json:"errors"`
|
||||
AvgLatencyMs float64 `gorm:"column:avg_latency_ms" json:"avgLatencyMs"`
|
||||
}
|
||||
|
||||
func (TunnelMetric) TableName() string { return "tunnel_metric" }
|
||||
|
||||
type ServiceMonitor struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
Name string `gorm:"type:varchar(100);not null" json:"name"`
|
||||
Type string `gorm:"type:varchar(20);not null" json:"type"`
|
||||
Target string `gorm:"type:text;not null" json:"target"`
|
||||
IntervalSec int `gorm:"column:interval_sec;not null;default:60" json:"intervalSec"`
|
||||
TimeoutSec int `gorm:"column:timeout_sec;not null;default:5" json:"timeoutSec"`
|
||||
NodeID int64 `gorm:"column:node_id;index" json:"nodeId"`
|
||||
Enabled int `gorm:"not null;default:1" json:"enabled"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null" json:"createdTime"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null" json:"updatedTime"`
|
||||
}
|
||||
|
||||
func (ServiceMonitor) TableName() string { return "service_monitor" }
|
||||
|
||||
type ServiceMonitorResult struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
MonitorID int64 `gorm:"column:monitor_id;not null;index:idx_monitor_result_monitor_time,priority:1" json:"monitorId"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;index" json:"nodeId"`
|
||||
Timestamp int64 `gorm:"not null;index:idx_monitor_result_monitor_time,priority:2" json:"timestamp"`
|
||||
Success int `gorm:"not null" json:"success"`
|
||||
LatencyMs float64 `gorm:"column:latency_ms" json:"latencyMs"`
|
||||
StatusCode int `gorm:"column:status_code" json:"statusCode"`
|
||||
ErrorMessage string `gorm:"column:error_message;type:text" json:"errorMessage"`
|
||||
}
|
||||
|
||||
func (ServiceMonitorResult) TableName() string { return "service_monitor_result" }
|
||||
|
||||
// TunnelQuality stores periodic probe results for a tunnel.
|
||||
// Unlike the old upsert model, rows accumulate for history/charting.
|
||||
// Old rows are pruned periodically (default: keep 24h).
|
||||
type TunnelQuality struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null;index:idx_tunnel_quality_tunnel_time,priority:1" json:"tunnelId"`
|
||||
EntryToExitLatency float64 `gorm:"column:entry_to_exit_latency" json:"entryToExitLatency"`
|
||||
ExitToBingLatency float64 `gorm:"column:exit_to_bing_latency" json:"exitToBingLatency"`
|
||||
EntryToExitLoss float64 `gorm:"column:entry_to_exit_loss" json:"entryToExitLoss"`
|
||||
ExitToBingLoss float64 `gorm:"column:exit_to_bing_loss" json:"exitToBingLoss"`
|
||||
Success int `gorm:"not null;default:1" json:"success"`
|
||||
ErrorMessage string `gorm:"column:error_message;type:text" json:"errorMessage,omitempty"`
|
||||
Timestamp int64 `gorm:"not null;index:idx_tunnel_quality_tunnel_time,priority:2;index:idx_tunnel_quality_time" json:"timestamp"`
|
||||
ChainDetails string `gorm:"column:chain_details;type:text" json:"chainDetails,omitempty"`
|
||||
}
|
||||
|
||||
func (TunnelQuality) TableName() string { return "tunnel_quality" }
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -30,8 +30,15 @@ func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
return r.ListForwardsByTunnelTx(r.db, tunnelID)
|
||||
}
|
||||
|
||||
func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]model.ForwardRecord, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
var forwards []model.Forward
|
||||
err := r.db.Where("tunnel_id = ?", tunnelID).Order("id ASC").Find(&forwards).Error
|
||||
err := tx.Where("tunnel_id = ?", tunnelID).Order("id ASC").Find(&forwards).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -46,6 +53,7 @@ func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -56,22 +64,96 @@ func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
|
||||
func (r *Repository) ListActiveTunnelIDsByNode(nodeID int64) ([]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var ids []int64
|
||||
err := r.db.Model(&model.ChainTunnel{}).
|
||||
Joins("JOIN tunnel ON tunnel.id = chain_tunnel.tunnel_id").
|
||||
Where("chain_tunnel.node_id = ? AND tunnel.status = 1", nodeID).
|
||||
Select("DISTINCT chain_tunnel.tunnel_id").
|
||||
Order("chain_tunnel.tunnel_id ASC").
|
||||
Pluck("chain_tunnel.tunnel_id", &ids).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListActiveForwardIDsByNode(nodeID int64) ([]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var ids []int64
|
||||
err := r.db.Model(&model.ForwardPort{}).
|
||||
Joins("JOIN forward ON forward.id = forward_port.forward_id").
|
||||
Where("forward_port.node_id = ? AND forward.status = 1", nodeID).
|
||||
Select("DISTINCT forward_port.forward_id").
|
||||
Order("forward_port.forward_id ASC").
|
||||
Pluck("forward_port.forward_id", &ids).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecord, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
return r.ListForwardPortsTx(r.db, forwardID)
|
||||
}
|
||||
|
||||
func (r *Repository) ListForwardPortsTx(tx *gorm.DB, forwardID int64) ([]model.ForwardPortRecord, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
var ports []model.ForwardPort
|
||||
err := r.db.Where("forward_id = ?", forwardID).Order("id ASC").Find(&ports).Error
|
||||
err := tx.Where("forward_id = ?", forwardID).Order("id ASC").Find(&ports).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows := make([]model.ForwardPortRecord, 0, len(ports))
|
||||
for _, p := range ports {
|
||||
rows = append(rows, model.ForwardPortRecord{NodeID: p.NodeID, Port: p.Port})
|
||||
inIP := ""
|
||||
if p.InIP.Valid {
|
||||
inIP = p.InIP.String
|
||||
}
|
||||
rows = append(rows, model.ForwardPortRecord{NodeID: p.NodeID, Port: p.Port, InIP: inIP})
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
|
||||
func (r *Repository) HasOtherForwardOnNodePort(nodeID int64, port int, currentForwardID int64) (bool, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return false, errors.New("repository not initialized")
|
||||
}
|
||||
return r.HasOtherForwardOnNodePortTx(r.db, nodeID, port, currentForwardID)
|
||||
}
|
||||
|
||||
func (r *Repository) HasOtherForwardOnNodePortTx(tx *gorm.DB, nodeID int64, port int, currentForwardID int64) (bool, error) {
|
||||
if tx == nil {
|
||||
return false, errors.New("database unavailable")
|
||||
}
|
||||
if nodeID <= 0 || port <= 0 {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
var count int64
|
||||
err := tx.Model(&model.ForwardPort{}).
|
||||
Where("node_id = ? AND port = ? AND forward_id <> ?", nodeID, port, currentForwardID).
|
||||
Count(&count).Error
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
|
||||
func (r *Repository) GetTunnelOutProtocol(tunnelID int64) (string, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return "", errors.New("repository not initialized")
|
||||
@@ -142,6 +224,9 @@ func nodeRecordFromModel(n *model.Node) *model.NodeRecord {
|
||||
if n.ServerIPV6.Valid {
|
||||
rec.ServerIPv6 = strings.TrimSpace(n.ServerIPV6.String)
|
||||
}
|
||||
if n.ExtraIPs.Valid {
|
||||
rec.ExtraIPs = strings.TrimSpace(n.ExtraIPs.String)
|
||||
}
|
||||
if n.InterfaceName.Valid {
|
||||
rec.InterfaceName = strings.TrimSpace(n.InterfaceName.String)
|
||||
}
|
||||
@@ -254,10 +339,11 @@ func (r *Repository) ListChainNodesForTunnel(tunnelID int64) ([]model.ChainNodeR
|
||||
Name sql.NullString
|
||||
Protocol sql.NullString
|
||||
Strategy sql.NullString
|
||||
ConnectIP sql.NullString
|
||||
}
|
||||
var rows []row
|
||||
err := r.db.Model(&model.ChainTunnel{}).
|
||||
Select("chain_tunnel.chain_type, chain_tunnel.inx, chain_tunnel.node_id, chain_tunnel.port, node.name, chain_tunnel.protocol, chain_tunnel.strategy").
|
||||
Select("chain_tunnel.chain_type, chain_tunnel.inx, chain_tunnel.node_id, chain_tunnel.port, node.name, chain_tunnel.protocol, chain_tunnel.strategy, chain_tunnel.connect_ip").
|
||||
Joins("LEFT JOIN node ON node.id = chain_tunnel.node_id").
|
||||
Where("chain_tunnel.tunnel_id = ?", tunnelID).
|
||||
Order("chain_tunnel.chain_type ASC, chain_tunnel.inx ASC, chain_tunnel.id ASC").
|
||||
@@ -302,6 +388,9 @@ func (r *Repository) ListChainNodesForTunnel(tunnelID int64) ([]model.ChainNodeR
|
||||
} else {
|
||||
item.Strategy = row.Strategy.String
|
||||
}
|
||||
if row.ConnectIP.Valid {
|
||||
item.ConnectIP = row.ConnectIP.String
|
||||
}
|
||||
result = append(result, item)
|
||||
}
|
||||
return result, nil
|
||||
|
||||
@@ -38,6 +38,14 @@ type FederationBindingRow struct {
|
||||
UpdatedTime int64
|
||||
}
|
||||
|
||||
type ActiveForwardPortRow struct {
|
||||
ForwardID int64
|
||||
TunnelID int64
|
||||
TunnelName string
|
||||
Port int
|
||||
UpdatedTime int64
|
||||
}
|
||||
|
||||
// ListRemoteNodes returns all nodes with is_remote=1, ordered by id desc.
|
||||
func (r *Repository) ListRemoteNodes() ([]RemoteNodeRow, error) {
|
||||
if r == nil || r.db == nil {
|
||||
@@ -87,6 +95,27 @@ func (r *Repository) ListActiveBindingsForNode(nodeID int64) ([]FederationBindin
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListActiveForwardPortsForNode(nodeID int64) ([]ActiveForwardPortRow, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var result []ActiveForwardPortRow
|
||||
err := r.db.Model(&model.ForwardPort{}).
|
||||
Select("forward_port.forward_id, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward_port.port, forward.updated_time").
|
||||
Joins("JOIN forward ON forward.id = forward_port.forward_id").
|
||||
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
|
||||
Where("forward_port.node_id = ? AND forward_port.port > 0", nodeID).
|
||||
Order("forward_port.port ASC, forward_port.id ASC").
|
||||
Find(&result).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if result == nil {
|
||||
result = make([]ActiveForwardPortRow, 0)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// GetNodeBasicInfo returns the name, server_ip, and status for a given node.
|
||||
func (r *Repository) GetNodeBasicInfo(nodeID int64) (*NodeBasicInfo, error) {
|
||||
if r == nil || r.db == nil {
|
||||
@@ -199,7 +228,6 @@ func (r *Repository) ListTunnelIDsByNamePrefix(prefix string) ([]int64, error) {
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
// NextIndex returns COALESCE(MAX(inx), -1) + 1 for the given table.
|
||||
func (r *Repository) NextIndex(table string) int {
|
||||
if r == nil || r.db == nil {
|
||||
return 0
|
||||
@@ -222,7 +250,7 @@ func (r *Repository) NextIndex(table string) int {
|
||||
var row inxRow
|
||||
err := r.db.Model(modelRef).
|
||||
Select("inx").
|
||||
Order("inx DESC").
|
||||
Order("inx ASC, id ASC").
|
||||
Limit(1).
|
||||
Take(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
@@ -231,10 +259,7 @@ func (r *Repository) NextIndex(table string) int {
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
if row.Inx < 0 {
|
||||
return 0
|
||||
}
|
||||
return row.Inx + 1
|
||||
return row.Inx - 1
|
||||
}
|
||||
|
||||
// CreateRemoteNode inserts a new remote node.
|
||||
|
||||
@@ -38,6 +38,7 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -68,6 +69,38 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
if strings.TrimSpace(rows[i].Strategy) == "" {
|
||||
rows[i].Strategy = "fifo"
|
||||
}
|
||||
}
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]model.ForwardRecord, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var forwards []model.Forward
|
||||
err := r.db.Where("user_id = ? AND tunnel_id = ?", userID, tunnelID).Order("id ASC").Find(&forwards).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -99,6 +132,7 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
}
|
||||
if strings.TrimSpace(fr.Strategy) == "" {
|
||||
fr.Strategy = "fifo"
|
||||
@@ -158,6 +192,64 @@ func (r *Repository) ForwardExists(forwardID int64) (bool, error) {
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// MapForwardIDsToTunnelIDs returns a mapping from forward.id to forward.tunnel_id.
|
||||
// Missing forward IDs are omitted from the returned map.
|
||||
func (r *Repository) MapForwardIDsToTunnelIDs(forwardIDs []int64) (map[int64]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if len(forwardIDs) == 0 {
|
||||
return map[int64]int64{}, nil
|
||||
}
|
||||
|
||||
// Deduplicate and filter invalid IDs.
|
||||
ids := make([]int64, 0, len(forwardIDs))
|
||||
seen := make(map[int64]struct{}, len(forwardIDs))
|
||||
for _, id := range forwardIDs {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
if len(ids) == 0 {
|
||||
return map[int64]int64{}, nil
|
||||
}
|
||||
|
||||
type row struct {
|
||||
ID int64 `gorm:"column:id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id"`
|
||||
}
|
||||
|
||||
out := make(map[int64]int64, len(ids))
|
||||
const chunkSize = 500
|
||||
for start := 0; start < len(ids); start += chunkSize {
|
||||
end := start + chunkSize
|
||||
if end > len(ids) {
|
||||
end = len(ids)
|
||||
}
|
||||
|
||||
var rows []row
|
||||
if err := r.db.Model(&model.Forward{}).
|
||||
Select("id", "tunnel_id").
|
||||
Where("id IN ?", ids[start:end]).
|
||||
Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, r := range rows {
|
||||
if r.ID <= 0 || r.TunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
out[r.ID] = r.TunnelID
|
||||
}
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Repository) SpeedLimitExists(id int64) (bool, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return false, errors.New("repository not initialized")
|
||||
@@ -169,3 +261,33 @@ func (r *Repository) SpeedLimitExists(id int64) (bool, error) {
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CountActiveForwardsByUser(userID int64) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
var count int64
|
||||
err := r.db.Model(&model.Forward{}).Where("user_id = ? AND status = 1", userID).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *Repository) CountActiveForwardsByUserTunnel(userID, tunnelID int64) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
var count int64
|
||||
err := r.db.Model(&model.Forward{}).Where("user_id = ? AND tunnel_id = ? AND status = 1", userID, tunnelID).Count(&count).Error
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (r *Repository) GetSpeedLimitSpeed(id int64) (int, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
var sl model.SpeedLimit
|
||||
err := r.db.Select("speed").Where("id = ?", id).First(&sl).Error
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return sl.Speed, nil
|
||||
}
|
||||
|
||||
@@ -50,8 +50,17 @@ func (r *Repository) ListGroupPermissionPairsByUserGroup(userGroupID int64) ([][
|
||||
return result, err
|
||||
}
|
||||
|
||||
// ListGroupPermissionPairsByTunnelGroup returns [userGroupID, tunnelGroupID] pairs
|
||||
// for all group permissions associated with a tunnel group.
|
||||
func (r *Repository) GetUserGroupIDsByUserID(userID int64) ([]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var ids []int64
|
||||
err := r.db.Model(&model.UserGroupUser{}).
|
||||
Where("user_id = ?", userID).
|
||||
Pluck("user_group_id", &ids).Error
|
||||
return ids, err
|
||||
}
|
||||
|
||||
func (r *Repository) ListGroupPermissionPairsByTunnelGroup(tunnelGroupID int64) ([][2]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
|
||||
@@ -1,14 +1,63 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
gsqlite "github.com/glebarez/sqlite"
|
||||
"go-backend/internal/store/model"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
func TestPrepareSQLiteLegacyColumnsAddsNodeMetadataColumns(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`
|
||||
CREATE TABLE node (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
secret VARCHAR(100) NOT NULL,
|
||||
server_ip VARCHAR(100) NOT NULL,
|
||||
port TEXT NOT NULL,
|
||||
interface_name VARCHAR(200),
|
||||
version VARCHAR(100),
|
||||
http INTEGER NOT NULL DEFAULT 0,
|
||||
tls INTEGER NOT NULL DEFAULT 0,
|
||||
socks INTEGER NOT NULL DEFAULT 0,
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER,
|
||||
status INTEGER NOT NULL
|
||||
)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create legacy node table: %v", err)
|
||||
}
|
||||
|
||||
if err := prepareSQLiteLegacyColumns(db); err != nil {
|
||||
t.Fatalf("prepareSQLiteLegacyColumns: %v", err)
|
||||
}
|
||||
|
||||
m := db.Migrator()
|
||||
for _, field := range []string{"Remark", "ExpiryTime", "RenewalCycle"} {
|
||||
if !m.HasColumn(&model.Node{}, field) {
|
||||
t.Fatalf("expected node.%s column to exist", field)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
@@ -83,3 +132,281 @@ func TestMigrateSchemaReturnsPostgresIDRepairError(t *testing.T) {
|
||||
t.Fatalf("expected error %v, got %v", wantErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaRunsViteConfigValueMigrationForLegacySchema(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, 2).Error; err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
|
||||
originalIDRepair := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = originalIDRepair
|
||||
})
|
||||
|
||||
called := 0
|
||||
originalMigrate := migrateViteConfigValueColumnTypeFn
|
||||
migrateViteConfigValueColumnTypeFn = func(db *gorm.DB) error {
|
||||
called++
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
migrateViteConfigValueColumnTypeFn = originalMigrate
|
||||
})
|
||||
|
||||
if err := migrateSchema(db); err != nil {
|
||||
t.Fatalf("migrateSchema: %v", err)
|
||||
}
|
||||
|
||||
if called != 1 {
|
||||
t.Fatalf("expected vite_config migration to run once, got %d", called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaReturnsViteConfigMigrationError(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, 2).Error; err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
|
||||
originalIDRepair := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = originalIDRepair
|
||||
})
|
||||
|
||||
wantErr := errors.New("vite config migration failed")
|
||||
originalMigrate := migrateViteConfigValueColumnTypeFn
|
||||
migrateViteConfigValueColumnTypeFn = func(db *gorm.DB) error {
|
||||
return wantErr
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
migrateViteConfigValueColumnTypeFn = originalMigrate
|
||||
})
|
||||
|
||||
err = migrateSchema(db)
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("expected error %v, got %v", wantErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaClearsSpeedLimitTunnelBinding(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, 3).Error; err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
if err := db.Exec(`
|
||||
CREATE TABLE speed_limit (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
speed INTEGER NOT NULL,
|
||||
tunnel_id INTEGER,
|
||||
tunnel_name VARCHAR(100),
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER,
|
||||
status INTEGER NOT NULL
|
||||
)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create speed_limit: %v", err)
|
||||
}
|
||||
if err := db.Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?)
|
||||
`, "legacy-speed-limit", 100, 101, "legacy-tunnel", 1, 1, 1).Error; err != nil {
|
||||
t.Fatalf("seed speed_limit: %v", err)
|
||||
}
|
||||
|
||||
originalIDRepair := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = originalIDRepair
|
||||
})
|
||||
|
||||
if err := migrateSchema(db); err != nil {
|
||||
t.Fatalf("migrateSchema: %v", err)
|
||||
}
|
||||
|
||||
var tunnelID sql.NullInt64
|
||||
var tunnelName sql.NullString
|
||||
if err := db.Raw(`SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?`, "legacy-speed-limit").Row().Scan(&tunnelID, &tunnelName); err != nil {
|
||||
t.Fatalf("query speed_limit: %v", err)
|
||||
}
|
||||
if tunnelID.Valid {
|
||||
t.Fatalf("expected tunnel_id cleared to NULL, got %d", tunnelID.Int64)
|
||||
}
|
||||
if tunnelName.Valid {
|
||||
t.Fatalf("expected tunnel_name cleared to NULL, got %q", tunnelName.String)
|
||||
}
|
||||
|
||||
var schemaVersion int
|
||||
if err := db.Raw(`SELECT version FROM schema_version LIMIT 1`).Row().Scan(&schemaVersion); err != nil {
|
||||
t.Fatalf("query schema_version: %v", err)
|
||||
}
|
||||
if schemaVersion != currentSchemaVersion {
|
||||
t.Fatalf("expected schema version %d, got %d", currentSchemaVersion, schemaVersion)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaRunsTrafficInt64MigrationForLegacySchema(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, 4).Error; err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
|
||||
originalIDRepair := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = originalIDRepair
|
||||
})
|
||||
|
||||
called := 0
|
||||
originalMigrate := migratePostgresTrafficInt64ColumnsFn
|
||||
migratePostgresTrafficInt64ColumnsFn = func(db *gorm.DB) error {
|
||||
called++
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
migratePostgresTrafficInt64ColumnsFn = originalMigrate
|
||||
})
|
||||
|
||||
if err := migrateSchema(db); err != nil {
|
||||
t.Fatalf("migrateSchema: %v", err)
|
||||
}
|
||||
|
||||
if called != 1 {
|
||||
t.Fatalf("expected traffic bigint migration to run once, got %d", called)
|
||||
}
|
||||
|
||||
var schemaVersion int
|
||||
if err := db.Raw(`SELECT version FROM schema_version LIMIT 1`).Row().Scan(&schemaVersion); err != nil {
|
||||
t.Fatalf("query schema_version: %v", err)
|
||||
}
|
||||
if schemaVersion != currentSchemaVersion {
|
||||
t.Fatalf("expected schema version %d, got %d", currentSchemaVersion, schemaVersion)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaReturnsTrafficInt64MigrationError(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
sqlDB, _ := db.DB()
|
||||
if sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
})
|
||||
|
||||
if err := db.Exec(`CREATE TABLE schema_version (version INTEGER NOT NULL DEFAULT 0)`).Error; err != nil {
|
||||
t.Fatalf("create schema_version: %v", err)
|
||||
}
|
||||
if err := db.Exec(`INSERT INTO schema_version(version) VALUES(?)`, 4).Error; err != nil {
|
||||
t.Fatalf("seed schema_version: %v", err)
|
||||
}
|
||||
|
||||
originalIDRepair := ensurePostgresIDDefaultsFn
|
||||
ensurePostgresIDDefaultsFn = func(db *gorm.DB) error {
|
||||
return nil
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
ensurePostgresIDDefaultsFn = originalIDRepair
|
||||
})
|
||||
|
||||
wantErr := errors.New("traffic bigint migration failed")
|
||||
originalMigrate := migratePostgresTrafficInt64ColumnsFn
|
||||
migratePostgresTrafficInt64ColumnsFn = func(db *gorm.DB) error {
|
||||
return wantErr
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
migratePostgresTrafficInt64ColumnsFn = originalMigrate
|
||||
})
|
||||
|
||||
err = migrateSchema(db)
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("expected error %v, got %v", wantErr, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAlterPostgresColumnToBigIntIfNeededValidatesNames(t *testing.T) {
|
||||
if err := alterPostgresColumnToBigIntIfNeeded(nil, "peer_share", "max_bandwidth"); err == nil || !strings.Contains(err.Error(), "nil db") {
|
||||
t.Fatalf("expected nil db error, got %v", err)
|
||||
}
|
||||
if err := alterPostgresColumnToBigIntIfNeeded(&gorm.DB{}, "", "max_bandwidth"); err == nil || !strings.Contains(err.Error(), "empty table or column name") {
|
||||
t.Fatalf("expected empty name error, got %v", err)
|
||||
}
|
||||
if err := alterPostgresColumnToBigIntIfNeeded(&gorm.DB{}, "peer_share", ""); err == nil || !strings.Contains(err.Error(), "empty table or column name") {
|
||||
t.Fatalf("expected empty name error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func (r *Repository) ListMonitorNodes() ([]model.Node, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var nodes []model.Node
|
||||
err := r.db.Select("id", "inx", "name", "status", "version", "updated_time").
|
||||
Where("is_remote = ?", 0).
|
||||
Order("inx ASC, id ASC").
|
||||
Find(&nodes).Error
|
||||
return nodes, err
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
func (r *Repository) InsertMonitorPermission(userID int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return nil
|
||||
}
|
||||
row := model.MonitorPermission{UserID: userID, CreatedTime: now}
|
||||
return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&row).Error
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteMonitorPermission(userID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return nil
|
||||
}
|
||||
return r.db.Where("user_id = ?", userID).Delete(&model.MonitorPermission{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) HasMonitorPermission(userID int64) (bool, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return false, errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return false, nil
|
||||
}
|
||||
var count int64
|
||||
err := r.db.Model(&model.MonitorPermission{}).Where("user_id = ?", userID).Count(&count).Error
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ListMonitorPermissions() ([]model.MonitorPermission, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var items []model.MonitorPermission
|
||||
err := r.db.Order("id ASC").Find(&items).Error
|
||||
return items, err
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func (r *Repository) ListMonitorTunnels() ([]model.Tunnel, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var tunnels []model.Tunnel
|
||||
err := r.db.Select("id", "inx", "name", "status", "updated_time").
|
||||
Order("inx ASC, id ASC").
|
||||
Find(&tunnels).Error
|
||||
return tunnels, err
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestGetTunnelMetricsAggregatedSumsAcrossNodes(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
ts := time.Now().UnixMilli()
|
||||
|
||||
if err := r.InsertTunnelMetric(&model.TunnelMetric{
|
||||
TunnelID: 1,
|
||||
NodeID: 1,
|
||||
Timestamp: ts,
|
||||
BytesIn: 100,
|
||||
BytesOut: 200,
|
||||
}); err != nil {
|
||||
t.Fatalf("insert tunnel metric n1: %v", err)
|
||||
}
|
||||
if err := r.InsertTunnelMetric(&model.TunnelMetric{
|
||||
TunnelID: 1,
|
||||
NodeID: 2,
|
||||
Timestamp: ts,
|
||||
BytesIn: 300,
|
||||
BytesOut: 400,
|
||||
}); err != nil {
|
||||
t.Fatalf("insert tunnel metric n2: %v", err)
|
||||
}
|
||||
|
||||
metrics, err := r.GetTunnelMetricsAggregated(1, ts-1000, ts+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get aggregated tunnel metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 {
|
||||
t.Fatalf("expected 1 aggregated point, got %d", len(metrics))
|
||||
}
|
||||
if metrics[0].Timestamp != ts {
|
||||
t.Fatalf("expected timestamp %d, got %d", ts, metrics[0].Timestamp)
|
||||
}
|
||||
if metrics[0].BytesIn != 400 {
|
||||
t.Fatalf("expected bytesIn 400, got %d", metrics[0].BytesIn)
|
||||
}
|
||||
if metrics[0].BytesOut != 600 {
|
||||
t.Fatalf("expected bytesOut 600, got %d", metrics[0].BytesOut)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpsertTunnelMetricBucketsAggregatesDuplicateKeysInBatch(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
ts := time.Now().UnixMilli()
|
||||
|
||||
items := []*model.TunnelMetric{
|
||||
{TunnelID: 1, NodeID: 1, Timestamp: ts, BytesIn: 10, BytesOut: 20},
|
||||
{TunnelID: 1, NodeID: 1, Timestamp: ts, BytesIn: 30, BytesOut: 40},
|
||||
}
|
||||
if err := r.UpsertTunnelMetricBuckets(items); err != nil {
|
||||
t.Fatalf("upsert buckets: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetTunnelMetrics(1, ts-1000, ts+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel metrics: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected 1 stored row, got %d", len(rows))
|
||||
}
|
||||
if rows[0].BytesIn != 40 {
|
||||
t.Fatalf("expected bytesIn 40, got %d", rows[0].BytesIn)
|
||||
}
|
||||
if rows[0].BytesOut != 60 {
|
||||
t.Fatalf("expected bytesOut 60, got %d", rows[0].BytesOut)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpsertTunnelMetricBucketsIsSafeUnderConcurrency(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
ts := time.Now().UnixMilli()
|
||||
|
||||
const workers = 20
|
||||
const perWorkerIn = int64(5)
|
||||
const perWorkerOut = int64(7)
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(workers)
|
||||
for i := 0; i < workers; i++ {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_ = r.UpsertTunnelMetricBuckets([]*model.TunnelMetric{{
|
||||
TunnelID: 1,
|
||||
NodeID: 1,
|
||||
Timestamp: ts,
|
||||
BytesIn: perWorkerIn,
|
||||
BytesOut: perWorkerOut,
|
||||
}})
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
rows, err := r.GetTunnelMetrics(1, ts-1000, ts+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel metrics: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected 1 stored row, got %d", len(rows))
|
||||
}
|
||||
|
||||
wantIn := int64(workers) * perWorkerIn
|
||||
wantOut := int64(workers) * perWorkerOut
|
||||
if rows[0].BytesIn != wantIn {
|
||||
t.Fatalf("expected bytesIn %d, got %d", wantIn, rows[0].BytesIn)
|
||||
}
|
||||
if rows[0].BytesOut != wantOut {
|
||||
t.Fatalf("expected bytesOut %d, got %d", wantOut, rows[0].BytesOut)
|
||||
}
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package repo
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -34,9 +35,9 @@ func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool
|
||||
return cnt > 0, err
|
||||
}
|
||||
|
||||
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status int, now int64) error {
|
||||
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status int, now int64) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
user := model.User{
|
||||
User: username,
|
||||
@@ -52,7 +53,10 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
}
|
||||
return r.db.Create(&user).Error
|
||||
if err := r.db.Create(&user).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return user.ID, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetUserRoleID(userID int64) (int, error) {
|
||||
@@ -141,6 +145,9 @@ func (r *Repository) DeleteUserCascade(userID int64) error {
|
||||
if err := tx.Where("user_id = ?", userID).Delete(&model.StatisticsFlow{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("user_id = ?", userID).Delete(&model.UserQuota{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Where("id = ?", userID).Delete(&model.User{}).Error
|
||||
})
|
||||
}
|
||||
@@ -193,16 +200,20 @@ 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 interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig 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{}) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
node := model.Node{
|
||||
Name: name,
|
||||
Remark: nullStringFromInterface(remark),
|
||||
ExpiryTime: nullInt64FromInterface(expiryTime),
|
||||
RenewalCycle: nullStringFromInterface(renewalCycle),
|
||||
Secret: secret,
|
||||
ServerIP: serverIP,
|
||||
ServerIPV4: nullStringFromInterface(serverIPV4),
|
||||
ServerIPV6: nullStringFromInterface(serverIPV6),
|
||||
ExtraIPs: nullStringFromInterface(extraIPs),
|
||||
Port: stringFromInterface(port),
|
||||
InterfaceName: nullStringFromInterface(interfaceName),
|
||||
Version: nullStringFromInterface(version),
|
||||
@@ -235,25 +246,30 @@ 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 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{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.Node{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"server_ip": serverIP,
|
||||
"server_ip_v4": nullStringFromInterface(serverIPV4),
|
||||
"server_ip_v6": nullStringFromInterface(serverIPV6),
|
||||
"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},
|
||||
"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
|
||||
}
|
||||
|
||||
@@ -293,6 +309,15 @@ func (r *Repository) UpdateNodeOrder(nodeID int64, inx int, now int64) {
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateNodeExpiryReminderDismissed(nodeID int64, dismissed int) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.Node{}).
|
||||
Where("id = ?", nodeID).
|
||||
Update("expiry_reminder_dismissed", dismissed).Error
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNodeCascade(nodeID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
@@ -392,7 +417,7 @@ func (r *Repository) DeleteChainTunnelsByTunnelTx(tx *gorm.DB, tunnelID int64) e
|
||||
return tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType string, nodeID int64, port sql.NullInt64, strategy string, inx int, protocol string) error {
|
||||
func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType string, nodeID int64, port sql.NullInt64, strategy string, inx int, protocol string, connectIp string) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
@@ -404,6 +429,7 @@ func (r *Repository) CreateChainTunnelTx(tx *gorm.DB, tunnelID int64, chainType
|
||||
Strategy: nullStringFromInterface(strategy),
|
||||
Inx: nullInt64FromInterface(inx),
|
||||
Protocol: nullStringFromInterface(protocol),
|
||||
ConnectIP: sql.NullString{String: connectIp, Valid: connectIp != ""},
|
||||
}
|
||||
return tx.Create(&ct).Error
|
||||
}
|
||||
@@ -519,9 +545,6 @@ func (r *Repository) DeleteTunnelCascade(tunnelID int64) error {
|
||||
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.SpeedLimit{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("tunnel_id = ?", tunnelID).Delete(&model.ChainTunnel{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -532,17 +555,6 @@ func (r *Repository) DeleteTunnelCascade(tunnelID int64) error {
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) GetTunnelNameByID(tunnelID int64) string {
|
||||
if r == nil || r.db == nil {
|
||||
return ""
|
||||
}
|
||||
var tunnel model.Tunnel
|
||||
if err := r.db.Select("name").Where("id = ?", tunnelID).First(&tunnel).Error; err != nil {
|
||||
return ""
|
||||
}
|
||||
return tunnel.Name
|
||||
}
|
||||
|
||||
func (r *Repository) TunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
@@ -654,7 +666,7 @@ 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) error {
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -665,6 +677,7 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
@@ -702,6 +715,7 @@ func (r *Repository) DeleteForwardCascade(forwardID int64) error {
|
||||
func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct {
|
||||
NodeID int64
|
||||
Port int
|
||||
InIP string
|
||||
}) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
@@ -715,13 +729,30 @@ func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct {
|
||||
}
|
||||
rows := make([]model.ForwardPort, 0, len(entries))
|
||||
for _, e := range entries {
|
||||
rows = append(rows, model.ForwardPort{ForwardID: forwardID, NodeID: e.NodeID, Port: e.Port})
|
||||
rows = append(rows, model.ForwardPort{
|
||||
ForwardID: forwardID,
|
||||
NodeID: e.NodeID,
|
||||
Port: e.Port,
|
||||
InIP: sql.NullString{String: e.InIP, Valid: e.InIP != ""},
|
||||
})
|
||||
}
|
||||
return tx.Create(&rows).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, now int64) {
|
||||
func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int, inIP string) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if forwardID <= 0 || nodeID <= 0 || port <= 0 {
|
||||
return nil
|
||||
}
|
||||
return r.db.Model(&model.ForwardPort{}).
|
||||
Where("forward_id = ? AND node_id = ? AND port = ?", forwardID, nodeID, port).
|
||||
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{}, now int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
@@ -735,6 +766,7 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"status": status,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
@@ -761,15 +793,15 @@ func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error)
|
||||
return used, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID int64, tunnelName string, now int64, status int) (int64, error) {
|
||||
func (r *Repository) CreateSpeedLimit(name string, speed int, now int64, status int) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
sl := model.SpeedLimit{
|
||||
Name: name,
|
||||
Speed: speed,
|
||||
TunnelID: tunnelID,
|
||||
TunnelName: tunnelName,
|
||||
TunnelID: sql.NullInt64{Int64: 0, Valid: false},
|
||||
TunnelName: sql.NullString{String: "", Valid: false},
|
||||
CreatedTime: now,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: status,
|
||||
@@ -780,34 +812,24 @@ func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID int64, tu
|
||||
return sl.ID, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID int64, tunnelName string, status int, now int64) error {
|
||||
func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, status int, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
updates := map[string]interface{}{
|
||||
"name": name,
|
||||
"speed": speed,
|
||||
"status": status,
|
||||
"tunnel_id": nil,
|
||||
"tunnel_name": nil,
|
||||
"updated_time": sql.NullInt64{
|
||||
Int64: now,
|
||||
Valid: true,
|
||||
},
|
||||
}
|
||||
return r.db.Model(&model.SpeedLimit{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"speed": speed,
|
||||
"tunnel_id": tunnelID,
|
||||
"tunnel_name": tunnelName,
|
||||
"status": status,
|
||||
"updated_time": sql.NullInt64{
|
||||
Int64: now,
|
||||
Valid: true,
|
||||
},
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) int64 {
|
||||
if r == nil || r.db == nil {
|
||||
return 0
|
||||
}
|
||||
var sl model.SpeedLimit
|
||||
if err := r.db.Select("tunnel_id").Where("id = ?", speedLimitID).First(&sl).Error; err != nil {
|
||||
return 0
|
||||
}
|
||||
return sl.TunnelID
|
||||
Updates(updates).Error
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteSpeedLimit(id int64) error {
|
||||
@@ -975,9 +997,16 @@ func (r *Repository) DeleteGroupPermissionByIDTx(tx *gorm.DB, id int64) error {
|
||||
return tx.Where("id = ?", id).Delete(&model.GroupPermission{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID int64, previousUserIDs, currentUserIDs []int64) error {
|
||||
// RevokedUserTunnelPair holds the (userID, tunnelID) of a deleted user_tunnel row,
|
||||
// so the handler layer can clean up associated forwarding rules.
|
||||
type RevokedUserTunnelPair struct {
|
||||
UserID int64
|
||||
TunnelID int64
|
||||
}
|
||||
|
||||
func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID int64, previousUserIDs, currentUserIDs []int64) ([]RevokedUserTunnelPair, error) {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
currentSet := make(map[int64]struct{}, len(currentUserIDs))
|
||||
for _, uid := range currentUserIDs {
|
||||
@@ -996,7 +1025,7 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID
|
||||
}
|
||||
}
|
||||
if len(removedUserIDs) == 0 {
|
||||
return nil
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
type grantRow struct {
|
||||
@@ -1004,6 +1033,8 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID
|
||||
CreatedByGroup int
|
||||
}
|
||||
|
||||
var revoked []RevokedUserTunnelPair
|
||||
|
||||
for _, userID := range removedUserIDs {
|
||||
var rows []grantRow
|
||||
if err := tx.Model(&model.GroupPermissionGrant{}).
|
||||
@@ -1011,7 +1042,7 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID
|
||||
Joins("JOIN user_tunnel ON user_tunnel.id = group_permission_grant.user_tunnel_id").
|
||||
Where("group_permission_grant.user_group_id = ? AND user_tunnel.user_id = ?", userGroupID, userID).
|
||||
Find(&rows).Error; err != nil {
|
||||
return err
|
||||
return revoked, err
|
||||
}
|
||||
|
||||
groupCreatedTunnelIDs := make(map[int64]struct{})
|
||||
@@ -1024,28 +1055,32 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID
|
||||
userTunnelIDs := tx.Model(&model.UserTunnel{}).Select("id").Where("user_id = ?", userID)
|
||||
if err := tx.Where("user_group_id = ? AND user_tunnel_id IN (?)", userGroupID, userTunnelIDs).
|
||||
Delete(&model.GroupPermissionGrant{}).Error; err != nil {
|
||||
return err
|
||||
return revoked, err
|
||||
}
|
||||
|
||||
for userTunnelID := range groupCreatedTunnelIDs {
|
||||
var remaining int64
|
||||
if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil {
|
||||
return err
|
||||
return revoked, err
|
||||
}
|
||||
if remaining == 0 {
|
||||
var ut model.UserTunnel
|
||||
if lookupErr := tx.Select("user_id", "tunnel_id").Where("id = ?", userTunnelID).First(&ut).Error; lookupErr == nil {
|
||||
revoked = append(revoked, RevokedUserTunnelPair{UserID: ut.UserID, TunnelID: ut.TunnelID})
|
||||
}
|
||||
if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
|
||||
return err
|
||||
return revoked, err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
return revoked, nil
|
||||
}
|
||||
|
||||
func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunnelGroupID int64) error {
|
||||
func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunnelGroupID int64) ([]RevokedUserTunnelPair, error) {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
|
||||
type grantRow struct {
|
||||
@@ -1058,7 +1093,7 @@ func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunne
|
||||
Select("user_tunnel_id, created_by_group").
|
||||
Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID).
|
||||
Find(&rows).Error; err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
groupCreatedTunnelIDs := make(map[int64]struct{})
|
||||
@@ -1070,22 +1105,27 @@ func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunne
|
||||
|
||||
if err := tx.Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID).
|
||||
Delete(&model.GroupPermissionGrant{}).Error; err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var revoked []RevokedUserTunnelPair
|
||||
for userTunnelID := range groupCreatedTunnelIDs {
|
||||
var remaining int64
|
||||
if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil {
|
||||
return err
|
||||
return revoked, err
|
||||
}
|
||||
if remaining == 0 {
|
||||
var ut model.UserTunnel
|
||||
if lookupErr := tx.Select("user_id", "tunnel_id").Where("id = ?", userTunnelID).First(&ut).Error; lookupErr == nil {
|
||||
revoked = append(revoked, RevokedUserTunnelPair{UserID: ut.UserID, TunnelID: ut.TunnelID})
|
||||
}
|
||||
if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil {
|
||||
return err
|
||||
return revoked, err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
return revoked, nil
|
||||
}
|
||||
|
||||
func (r *Repository) ReplaceFederationTunnelBindingsTx(tx *gorm.DB, tunnelID int64, bindings []FederationTunnelBinding) error {
|
||||
@@ -1187,7 +1227,7 @@ 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) (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{}) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
@@ -1206,6 +1246,7 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
}
|
||||
if err := tx.Create(&fwd).Error; err != nil {
|
||||
return err
|
||||
@@ -1216,6 +1257,7 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
|
||||
ForwardID: forwardID,
|
||||
NodeID: nodeID,
|
||||
Port: port,
|
||||
InIP: sql.NullString{String: inIp, Valid: inIp != ""},
|
||||
}
|
||||
if err := tx.Create(&fp).Error; err != nil {
|
||||
return err
|
||||
@@ -1406,3 +1448,126 @@ func parsePortRangeSpec(input string) []int {
|
||||
sort.Ints(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func (r *Repository) AddUserToGroups(userID int64, groupIDs []int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if len(groupIDs) == 0 {
|
||||
return nil
|
||||
}
|
||||
rows := make([]model.UserGroupUser, 0, len(groupIDs))
|
||||
for _, gid := range groupIDs {
|
||||
if gid <= 0 {
|
||||
continue
|
||||
}
|
||||
rows = append(rows, model.UserGroupUser{UserGroupID: gid, UserID: userID, CreatedTime: now})
|
||||
}
|
||||
if len(rows) == 0 {
|
||||
return nil
|
||||
}
|
||||
return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ReplaceUserGroupsByUserID(userID int64, newGroupIDs []int64, now int64) (affectedGroupIDs []int64, err error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
var oldGroupIDs []int64
|
||||
if err = r.db.Model(&model.UserGroupUser{}).
|
||||
Where("user_id = ?", userID).
|
||||
Pluck("user_group_id", &oldGroupIDs).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
seen := make(map[int64]struct{})
|
||||
for _, id := range oldGroupIDs {
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
for _, id := range newGroupIDs {
|
||||
if id > 0 {
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
for id := range seen {
|
||||
affectedGroupIDs = append(affectedGroupIDs, id)
|
||||
}
|
||||
|
||||
if err = r.db.Where("user_id = ?", userID).Delete(&model.UserGroupUser{}).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if len(newGroupIDs) == 0 {
|
||||
return affectedGroupIDs, nil
|
||||
}
|
||||
rows := make([]model.UserGroupUser, 0, len(newGroupIDs))
|
||||
for _, gid := range newGroupIDs {
|
||||
if gid <= 0 {
|
||||
continue
|
||||
}
|
||||
rows = append(rows, model.UserGroupUser{UserGroupID: gid, UserID: userID, CreatedTime: now})
|
||||
}
|
||||
if len(rows) > 0 {
|
||||
if err = r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return affectedGroupIDs, nil
|
||||
}
|
||||
|
||||
func (r *Repository) AdvanceNodeRenewalCycles(now int64) (int, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var nodes []model.Node
|
||||
if err := r.db.Where("renewal_cycle IS NOT NULL AND renewal_cycle != '' AND expiry_time IS NOT NULL").Find(&nodes).Error; err != nil {
|
||||
return 0, fmt.Errorf("list nodes with renewal cycle: %w", err)
|
||||
}
|
||||
|
||||
advanced := 0
|
||||
for _, node := range nodes {
|
||||
if !node.ExpiryTime.Valid || node.ExpiryTime.Int64 <= 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
cycleMonths := 0
|
||||
switch node.RenewalCycle.String {
|
||||
case "month":
|
||||
cycleMonths = 1
|
||||
case "quarter":
|
||||
cycleMonths = 3
|
||||
case "year":
|
||||
cycleMonths = 12
|
||||
default:
|
||||
continue
|
||||
}
|
||||
|
||||
anchorTime := node.ExpiryTime.Int64
|
||||
for anchorTime <= now {
|
||||
nextAnchor := advanceByMonths(anchorTime, cycleMonths)
|
||||
if nextAnchor <= anchorTime {
|
||||
break
|
||||
}
|
||||
anchorTime = nextAnchor
|
||||
}
|
||||
|
||||
if anchorTime == node.ExpiryTime.Int64 {
|
||||
continue
|
||||
}
|
||||
|
||||
if err := r.db.Model(&model.Node{}).Where("id = ?", node.ID).Update("expiry_time", anchorTime).Error; err != nil {
|
||||
continue
|
||||
}
|
||||
advanced++
|
||||
}
|
||||
|
||||
return advanced, nil
|
||||
}
|
||||
|
||||
func advanceByMonths(timestamp int64, months int) int64 {
|
||||
t := time.Unix(timestamp/1000, 0)
|
||||
next := t.AddDate(0, months, 0)
|
||||
return next.UnixMilli()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
// InsertTunnelQuality appends a tunnel quality probe result.
|
||||
// (Follows the same pattern as InsertServiceMonitorResult.)
|
||||
func (r *Repository) InsertTunnelQuality(q *model.TunnelQuality) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if q == nil || q.TunnelID <= 0 {
|
||||
return nil
|
||||
}
|
||||
return r.db.Create(q).Error
|
||||
}
|
||||
|
||||
// GetTunnelQualityHistory returns quality probe results for a tunnel
|
||||
// within a time range, ordered by timestamp ascending.
|
||||
// (Mirrors GetServiceMonitorResults pattern.)
|
||||
func (r *Repository) GetTunnelQualityHistory(tunnelID int64, startMs, endMs int64) ([]model.TunnelQuality, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var results []model.TunnelQuality
|
||||
err := r.db.Where("tunnel_id = ? AND timestamp >= ? AND timestamp <= ?", tunnelID, startMs, endMs).
|
||||
Order("timestamp ASC").
|
||||
Find(&results).Error
|
||||
return results, err
|
||||
}
|
||||
|
||||
// GetLatestTunnelQualities returns the newest quality result per tunnel_id.
|
||||
// (Mirrors GetLatestServiceMonitorResults pattern.)
|
||||
func (r *Repository) GetLatestTunnelQualities() ([]model.TunnelQuality, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var results []model.TunnelQuality
|
||||
|
||||
// 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
|
||||
FROM (
|
||||
SELECT *, ROW_NUMBER() OVER (PARTITION BY tunnel_id ORDER BY timestamp DESC, id DESC) AS rn
|
||||
FROM tunnel_quality
|
||||
) t
|
||||
WHERE rn = 1
|
||||
ORDER BY tunnel_id ASC
|
||||
`
|
||||
if err := r.db.Raw(q).Scan(&results).Error; err == nil {
|
||||
return results, nil
|
||||
}
|
||||
|
||||
// Fallback for older SQLite
|
||||
results = nil
|
||||
err := r.db.Order("timestamp DESC, id DESC").Limit(5000).Find(&results).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
seen := make(map[int64]struct{}, len(results))
|
||||
out := make([]model.TunnelQuality, 0, len(results))
|
||||
for _, row := range results {
|
||||
if row.TunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[row.TunnelID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[row.TunnelID] = struct{}{}
|
||||
out = append(out, row)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// PruneTunnelQualityResults deletes quality results older than the given timestamp.
|
||||
// (Mirrors PruneServiceMonitorResults pattern.)
|
||||
func (r *Repository) PruneTunnelQualityResults(olderThanMs int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return nil
|
||||
}
|
||||
return r.db.Where("timestamp < ?", olderThanMs).Delete(&model.TunnelQuality{}).Error
|
||||
}
|
||||
|
||||
// ListEnabledTunnelIDs returns IDs of all tunnels with status=1.
|
||||
func (r *Repository) ListEnabledTunnelIDs() ([]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var ids []int64
|
||||
err := r.db.Model(&model.Tunnel{}).Where("status = ?", 1).Pluck("id", &ids).Error
|
||||
return ids, err
|
||||
}
|
||||
@@ -0,0 +1,378 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const userQuotaBytesPerGB int64 = 1024 * 1024 * 1024
|
||||
|
||||
type UserQuotaRelease struct {
|
||||
UserID int64
|
||||
ForwardIDs []int64
|
||||
UnblockUser bool
|
||||
}
|
||||
|
||||
func userQuotaWindowKeys(now time.Time) (int64, int64) {
|
||||
return int64(now.Year()*10000 + int(now.Month())*100 + now.Day()), int64(now.Year()*100 + int(now.Month()))
|
||||
}
|
||||
|
||||
func cloneUserQuotaView(q model.UserQuota) *model.UserQuotaView {
|
||||
return &model.UserQuotaView{
|
||||
UserID: q.UserID,
|
||||
DailyLimitGB: q.DailyLimitGB,
|
||||
MonthlyLimitGB: q.MonthlyLimitGB,
|
||||
DailyUsedBytes: q.DailyUsedBytes,
|
||||
MonthlyUsedBytes: q.MonthlyUsedBytes,
|
||||
DayKey: q.DayKey,
|
||||
MonthKey: q.MonthKey,
|
||||
DisabledByQuota: q.DisabledByQuota,
|
||||
DisabledAt: q.DisabledAt,
|
||||
PausedForwardIDs: q.PausedForwardIDs,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeUserQuotaView(view *model.UserQuotaView, now time.Time) *model.UserQuotaView {
|
||||
if view == nil {
|
||||
return nil
|
||||
}
|
||||
dayKey, monthKey := userQuotaWindowKeys(now)
|
||||
out := *view
|
||||
if out.DayKey != dayKey {
|
||||
out.DayKey = dayKey
|
||||
out.DailyUsedBytes = 0
|
||||
}
|
||||
if out.MonthKey != monthKey {
|
||||
out.MonthKey = monthKey
|
||||
out.MonthlyUsedBytes = 0
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
func userQuotaExceeded(view *model.UserQuotaView) bool {
|
||||
if view == nil {
|
||||
return false
|
||||
}
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*userQuotaBytesPerGB {
|
||||
return true
|
||||
}
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*userQuotaBytesPerGB {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func parsePausedForwardIDs(raw string) []int64 {
|
||||
parts := strings.Split(strings.TrimSpace(raw), ",")
|
||||
out := make([]int64, 0, len(parts))
|
||||
seen := make(map[int64]struct{}, len(parts))
|
||||
for _, part := range parts {
|
||||
id, err := strconv.ParseInt(strings.TrimSpace(part), 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
out = append(out, id)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func joinPausedForwardIDs(ids []int64) string {
|
||||
if len(ids) == 0 {
|
||||
return ""
|
||||
}
|
||||
parts := make([]string, 0, len(ids))
|
||||
seen := make(map[int64]struct{}, len(ids))
|
||||
for _, id := range ids {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
parts = append(parts, strconv.FormatInt(id, 10))
|
||||
}
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
func (r *Repository) loadOrCreateUserQuotaTx(tx *gorm.DB, userID int64, now time.Time) (*model.UserQuota, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
dayKey, monthKey := userQuotaWindowKeys(now)
|
||||
q := &model.UserQuota{}
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("user_id = ?", userID).First(q).Error
|
||||
if err == nil {
|
||||
return q, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
q = &model.UserQuota{
|
||||
UserID: userID,
|
||||
DayKey: dayKey,
|
||||
MonthKey: monthKey,
|
||||
CreatedTime: nowMs,
|
||||
UpdatedTime: nowMs,
|
||||
PausedForwardIDs: "",
|
||||
}
|
||||
if err := tx.Create(q).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q, nil
|
||||
}
|
||||
|
||||
func applyUserQuotaWindowRoll(q *model.UserQuota, now time.Time) bool {
|
||||
if q == nil {
|
||||
return false
|
||||
}
|
||||
changed := false
|
||||
dayKey, monthKey := userQuotaWindowKeys(now)
|
||||
if q.DayKey != dayKey {
|
||||
q.DayKey = dayKey
|
||||
q.DailyUsedBytes = 0
|
||||
changed = true
|
||||
}
|
||||
if q.MonthKey != monthKey {
|
||||
q.MonthKey = monthKey
|
||||
q.MonthlyUsedBytes = 0
|
||||
changed = true
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
func (r *Repository) SaveUserQuotaConfigTx(tx *gorm.DB, userID, dailyLimitGB, monthlyLimitGB int64, now int64) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return errors.New("user id is required")
|
||||
}
|
||||
if dailyLimitGB < 0 || monthlyLimitGB < 0 {
|
||||
return errors.New("quota limit cannot be negative")
|
||||
}
|
||||
current := time.UnixMilli(now)
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, current)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
updates := map[string]interface{}{
|
||||
"daily_limit_gb": dailyLimitGB,
|
||||
"monthly_limit_gb": monthlyLimitGB,
|
||||
"updated_time": now,
|
||||
}
|
||||
if q.DayKey == 0 || q.MonthKey == 0 {
|
||||
dayKey, monthKey := userQuotaWindowKeys(current)
|
||||
updates["day_key"] = dayKey
|
||||
updates["month_key"] = monthKey
|
||||
}
|
||||
return tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(updates).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListUserQuotaViewsByUserIDs(userIDs []int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
out := make(map[int64]*model.UserQuotaView)
|
||||
if len(userIDs) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
var rows []model.UserQuota
|
||||
if err := r.db.Where("user_id IN ?", userIDs).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, row := range rows {
|
||||
out[row.UserID] = normalizeUserQuotaView(cloneUserQuotaView(row), now)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetUserQuotaView(userID int64, now time.Time) (*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var row model.UserQuota
|
||||
err := r.db.Where("user_id = ?", userID).First(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeUserQuotaView(cloneUserQuotaView(row), now), nil
|
||||
}
|
||||
|
||||
func (r *Repository) AddUserQuotaUsage(userID int64, usedBytes int64, now time.Time) (*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
result := &model.UserQuotaView{}
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usedBytes > 0 {
|
||||
q.DailyUsedBytes += usedBytes
|
||||
q.MonthlyUsedBytes += usedBytes
|
||||
}
|
||||
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 = *cloneUserQuotaView(*q)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeUserQuotaView(result, now), nil
|
||||
}
|
||||
|
||||
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return errors.New("user id is required")
|
||||
}
|
||||
return r.db.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
|
||||
"disabled_by_quota": 1,
|
||||
"disabled_at": now,
|
||||
"paused_forward_ids": joinPausedForwardIDs(pausedForwardIDs),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ResetUserQuotaUsage(userID int64, scope string, now time.Time) (*UserQuotaRelease, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if userID <= 0 {
|
||||
return nil, errors.New("user id is required")
|
||||
}
|
||||
scope = strings.TrimSpace(strings.ToLower(scope))
|
||||
if scope == "" {
|
||||
scope = "all"
|
||||
}
|
||||
if scope != "daily" && scope != "monthly" && scope != "all" {
|
||||
return nil, fmt.Errorf("unsupported quota reset scope: %s", scope)
|
||||
}
|
||||
var release *UserQuotaRelease
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
switch scope {
|
||||
case "daily":
|
||||
q.DailyUsedBytes = 0
|
||||
case "monthly":
|
||||
q.MonthlyUsedBytes = 0
|
||||
case "all":
|
||||
q.DailyUsedBytes = 0
|
||||
q.MonthlyUsedBytes = 0
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
release = &UserQuotaRelease{UserID: userID}
|
||||
if q.DisabledByQuota == 1 && !userQuotaExceeded(cloneUserQuotaView(*q)) {
|
||||
release.UnblockUser = true
|
||||
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
|
||||
q.DisabledByQuota = 0
|
||||
q.DisabledAt = 0
|
||||
q.PausedForwardIDs = ""
|
||||
}
|
||||
return 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,
|
||||
"disabled_by_quota": q.DisabledByQuota,
|
||||
"disabled_at": q.DisabledAt,
|
||||
"paused_forward_ids": q.PausedForwardIDs,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return release, nil
|
||||
}
|
||||
|
||||
func (r *Repository) RollUserQuotaWindows(now time.Time) ([]UserQuotaRelease, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var releases []UserQuotaRelease
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
var rows []model.UserQuota
|
||||
if err := tx.Find(&rows).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
for _, row := range rows {
|
||||
q := row
|
||||
changed := applyUserQuotaWindowRoll(&q, now)
|
||||
release := UserQuotaRelease{UserID: q.UserID}
|
||||
if q.DisabledByQuota == 1 && !userQuotaExceeded(cloneUserQuotaView(q)) {
|
||||
release.UnblockUser = true
|
||||
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
|
||||
q.DisabledByQuota = 0
|
||||
q.DisabledAt = 0
|
||||
q.PausedForwardIDs = ""
|
||||
changed = true
|
||||
}
|
||||
if !changed {
|
||||
continue
|
||||
}
|
||||
q.UpdatedTime = nowMs
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", q.UserID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"disabled_by_quota": q.DisabledByQuota,
|
||||
"disabled_at": q.DisabledAt,
|
||||
"paused_forward_ids": q.PausedForwardIDs,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if release.UnblockUser {
|
||||
releases = append(releases, release)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return releases, nil
|
||||
}
|
||||
@@ -39,6 +39,7 @@ type nodeSession struct {
|
||||
nodeID int64
|
||||
secret string
|
||||
conn *connWrap
|
||||
crypto *security.AESCrypto // 缓存的 AES 加密器,避免每条消息重建
|
||||
}
|
||||
|
||||
type commandResponse struct {
|
||||
@@ -68,9 +69,11 @@ type CommandResult struct {
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
upgrader websocket.Upgrader
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
upgrader websocket.Upgrader
|
||||
onNodeOnline func(nodeID int64)
|
||||
onNodeMetric func(nodeID int64, info SystemInfo)
|
||||
|
||||
mu sync.RWMutex
|
||||
admins map[*connWrap]struct{}
|
||||
@@ -79,6 +82,40 @@ type Server struct {
|
||||
pending map[string]pendingRequest
|
||||
}
|
||||
|
||||
type SystemInfo struct {
|
||||
Uptime uint64 `json:"uptime"`
|
||||
BytesReceived uint64 `json:"bytes_received"`
|
||||
BytesTransmitted uint64 `json:"bytes_transmitted"`
|
||||
CPUUsage float64 `json:"cpu_usage"`
|
||||
MemoryUsage float64 `json:"memory_usage"`
|
||||
DiskUsage float64 `json:"disk_usage"`
|
||||
Load1 float64 `json:"load1"`
|
||||
Load5 float64 `json:"load5"`
|
||||
Load15 float64 `json:"load15"`
|
||||
TCPConns int64 `json:"tcp_conns"`
|
||||
UDPConns int64 `json:"udp_conns"`
|
||||
NetInSpeed int64 `json:"net_in_speed"`
|
||||
NetOutSpeed int64 `json:"net_out_speed"`
|
||||
}
|
||||
|
||||
func (s *Server) SetNodeOnlineHook(fn func(nodeID int64)) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.onNodeOnline = fn
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func (s *Server) SetNodeMetricHook(fn func(nodeID int64, info SystemInfo)) {
|
||||
if s == nil {
|
||||
return
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.onNodeMetric = fn
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
func NewServer(repo *repo.Repository, jwtSecret string) *Server {
|
||||
return &Server{
|
||||
repo: repo,
|
||||
@@ -175,7 +212,12 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
_ = old.conn.conn.Close()
|
||||
delete(s.byConn, old.conn.conn)
|
||||
}
|
||||
ns := &nodeSession{nodeID: nodeID, secret: secret, conn: cw}
|
||||
// 初始化 AES 加密器并缓存(仅创建一次)
|
||||
var nodeCrypto *security.AESCrypto
|
||||
if strings.TrimSpace(secret) != "" {
|
||||
nodeCrypto, _ = security.NewAESCrypto(secret)
|
||||
}
|
||||
ns := &nodeSession{nodeID: nodeID, secret: secret, conn: cw, crypto: nodeCrypto}
|
||||
s.nodes[nodeID] = ns
|
||||
s.byConn[conn] = ns
|
||||
s.mu.Unlock()
|
||||
@@ -183,6 +225,13 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
_ = s.repo.UpdateNodeOnline(nodeID, 1, version, httpVal, tlsVal, socksVal)
|
||||
s.broadcastStatus(nodeID, 1)
|
||||
|
||||
s.mu.RLock()
|
||||
onlineHook := s.onNodeOnline
|
||||
s.mu.RUnlock()
|
||||
if onlineHook != nil {
|
||||
go onlineHook(nodeID)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
close(done)
|
||||
needOfflineBroadcast := false
|
||||
@@ -208,18 +257,100 @@ func (s *Server) handleNode(w http.ResponseWriter, r *http.Request, nodeID int64
|
||||
return
|
||||
}
|
||||
|
||||
msg := decryptIfNeeded(payload, secret)
|
||||
msg := decryptIfNeeded(payload, ns.crypto, secret)
|
||||
s.tryResolvePending(nodeID, msg)
|
||||
|
||||
var parsed struct {
|
||||
Type string `json:"type"`
|
||||
}
|
||||
if json.Unmarshal([]byte(msg), &parsed) == nil && parsed.Type == "UpgradeProgress" {
|
||||
s.broadcastTyped(nodeID, "upgrade_progress", msg)
|
||||
} else {
|
||||
s.broadcastInfo(nodeID, msg)
|
||||
if json.Unmarshal([]byte(msg), &parsed) == nil && parsed.Type != "" {
|
||||
switch parsed.Type {
|
||||
case "metric":
|
||||
// Agent 新版指标消息:{type:"metric", data:{...}}
|
||||
var envelope struct {
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(msg), &envelope); err == nil && len(envelope.Data) > 0 {
|
||||
// 解析 SystemInfo 并调用 hook
|
||||
var sysInfo SystemInfo
|
||||
if json.Unmarshal(envelope.Data, &sysInfo) == nil {
|
||||
s.mu.RLock()
|
||||
onMetric := s.onNodeMetric
|
||||
s.mu.RUnlock()
|
||||
if onMetric != nil {
|
||||
go onMetric(nodeID, sysInfo)
|
||||
}
|
||||
}
|
||||
// 广播内层 data 给前端(保持平坦结构兼容性)
|
||||
s.broadcastTyped(nodeID, "metric", string(envelope.Data))
|
||||
}
|
||||
continue
|
||||
case "UpgradeProgress":
|
||||
s.broadcastTyped(nodeID, "upgrade_progress", msg)
|
||||
continue
|
||||
default:
|
||||
// Unknown typed messages still get broadcast so future
|
||||
// agent message types are not silently lost.
|
||||
s.broadcastInfo(nodeID, msg)
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
// 兼容旧版 Agent:无 type 字段的系统信息消息
|
||||
if looksLikeSystemInfoMessage(msg) {
|
||||
var sysInfo SystemInfo
|
||||
if err := json.Unmarshal([]byte(msg), &sysInfo); err == nil {
|
||||
s.mu.RLock()
|
||||
onMetric := s.onNodeMetric
|
||||
s.mu.RUnlock()
|
||||
if onMetric != nil {
|
||||
go onMetric(nodeID, sysInfo)
|
||||
}
|
||||
s.broadcastTyped(nodeID, "metric", msg)
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
s.broadcastInfo(nodeID, msg)
|
||||
}
|
||||
}
|
||||
|
||||
func looksLikeSystemInfoMessage(msg string) bool {
|
||||
// Keep this as a cheap heuristic so that arbitrary JSON objects don't get
|
||||
// misclassified as metrics (SystemInfo unmarshal would otherwise succeed with
|
||||
// all-zero values).
|
||||
if strings.TrimSpace(msg) == "" {
|
||||
return false
|
||||
}
|
||||
if !strings.Contains(msg, "{") {
|
||||
return false
|
||||
}
|
||||
|
||||
keys := []string{
|
||||
"\"uptime\"",
|
||||
"\"cpu_usage\"",
|
||||
"\"memory_usage\"",
|
||||
"\"disk_usage\"",
|
||||
"\"bytes_received\"",
|
||||
"\"bytes_transmitted\"",
|
||||
"\"net_in_speed\"",
|
||||
"\"net_out_speed\"",
|
||||
"\"tcp_conns\"",
|
||||
"\"udp_conns\"",
|
||||
"\"load1\"",
|
||||
"\"load5\"",
|
||||
"\"load15\"",
|
||||
}
|
||||
matched := 0
|
||||
for _, k := range keys {
|
||||
if strings.Contains(msg, k) {
|
||||
matched++
|
||||
if matched >= 3 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, timeout time.Duration) (CommandResult, error) {
|
||||
@@ -268,13 +399,8 @@ func (s *Server) SendCommand(nodeID int64, cmdType string, data interface{}, tim
|
||||
}
|
||||
|
||||
messageData := rawCmd
|
||||
if strings.TrimSpace(ns.secret) != "" {
|
||||
crypto, err := security.NewAESCrypto(ns.secret)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return CommandResult{}, err
|
||||
}
|
||||
encrypted, err := crypto.Encrypt(rawCmd)
|
||||
if ns.crypto != nil {
|
||||
encrypted, err := ns.crypto.Encrypt(rawCmd)
|
||||
if err != nil {
|
||||
cleanup()
|
||||
return CommandResult{}, err
|
||||
@@ -324,6 +450,11 @@ func (s *Server) tryResolvePending(nodeID int64, message string) {
|
||||
return
|
||||
}
|
||||
|
||||
// 快速短路:指标消息永远不含 requestId,跳过完整 JSON 解析
|
||||
if !strings.Contains(message, "\"requestId\"") {
|
||||
return
|
||||
}
|
||||
|
||||
var resp commandResponse
|
||||
if err := json.Unmarshal([]byte(message), &resp); err != nil {
|
||||
return
|
||||
@@ -441,18 +572,22 @@ func (s *Server) broadcastToAdmins(message string) {
|
||||
}
|
||||
}
|
||||
|
||||
func decryptIfNeeded(payload []byte, secret string) string {
|
||||
func decryptIfNeeded(payload []byte, crypto *security.AESCrypto, secret string) string {
|
||||
text := string(payload)
|
||||
var wrap encryptedMessage
|
||||
if err := json.Unmarshal(payload, &wrap); err != nil || !wrap.Encrypted || strings.TrimSpace(wrap.Data) == "" {
|
||||
return text
|
||||
}
|
||||
|
||||
crypto, err := security.NewAESCrypto(secret)
|
||||
if err != nil {
|
||||
// 优先使用缓存的 crypto 实例
|
||||
c := crypto
|
||||
if c == nil && strings.TrimSpace(secret) != "" {
|
||||
c, _ = security.NewAESCrypto(secret)
|
||||
}
|
||||
if c == nil {
|
||||
return text
|
||||
}
|
||||
plain, err := crypto.Decrypt(wrap.Data)
|
||||
plain, err := c.Decrypt(wrap.Data)
|
||||
if err != nil {
|
||||
return text
|
||||
}
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestForwardBatchDeleteReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-delete", `{"ids":[999]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureReasonContains(t, result, "转发不存在")
|
||||
}
|
||||
|
||||
func TestForwardBatchPauseReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-pause", `{"ids":[999]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureReasonContains(t, result, "转发不存在")
|
||||
}
|
||||
|
||||
func TestForwardBatchResumeReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
forwardID := seedForwardForBatchAction(t, repo, batchForwardSeedOptions{
|
||||
Now: now,
|
||||
TunnelName: "resume-detail-tunnel",
|
||||
ForwardName: "resume-detail-forward",
|
||||
CreateUserTunnel: true,
|
||||
UserTunnelStatus: 0,
|
||||
})
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-resume", `{"ids":[`+jsonNumber(forwardID)+`]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureNameAndReason(t, result, "resume-detail-forward", "该隧道已禁用")
|
||||
}
|
||||
|
||||
func TestForwardBatchChangeTunnelReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
forwardID := seedForwardForBatchAction(t, repo, batchForwardSeedOptions{
|
||||
Now: now,
|
||||
TunnelName: "change-detail-tunnel",
|
||||
ForwardName: "change-detail-forward",
|
||||
})
|
||||
tunnelID := mustQueryInt64(t, repo, `SELECT tunnel_id FROM forward WHERE id = ?`, forwardID)
|
||||
|
||||
payload := `{"forwardIds":[` + jsonNumber(forwardID) + `],"targetTunnelId":` + jsonNumber(tunnelID) + `}`
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/forward/batch-change-tunnel", payload)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureNameAndReason(t, result, "change-detail-forward", "规则已在目标隧道中")
|
||||
}
|
||||
|
||||
func TestTunnelBatchDeleteReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
|
||||
out := postBatchRequest(t, router, adminToken, "/api/v1/tunnel/batch-delete", `{"ids":[999]}`)
|
||||
result := mustBatchResult(t, out)
|
||||
assertBatchFailureReasonContains(t, result, "隧道不存在")
|
||||
}
|
||||
|
||||
type batchForwardSeedOptions struct {
|
||||
Now int64
|
||||
TunnelName string
|
||||
ForwardName string
|
||||
CreateUserTunnel bool
|
||||
UserTunnelStatus int
|
||||
}
|
||||
|
||||
func mustAdminToken(t *testing.T, secret string) string {
|
||||
t.Helper()
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
return token
|
||||
}
|
||||
|
||||
func postBatchRequest(t *testing.T, router http.Handler, token, path, payload string) response.R {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected API success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func mustBatchResult(t *testing.T, out response.R) map[string]interface{} {
|
||||
t.Helper()
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected map result, got %T", out.Data)
|
||||
}
|
||||
if int(result["failCount"].(float64)) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %v", result["failCount"])
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func assertBatchFailureReasonContains(t *testing.T, result map[string]interface{}, snippet string) {
|
||||
t.Helper()
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, snippet) {
|
||||
t.Fatalf("expected failure reason to contain %q, got %q", snippet, reason)
|
||||
}
|
||||
}
|
||||
|
||||
func assertBatchFailureNameAndReason(t *testing.T, result map[string]interface{}, expectedName, reasonSnippet string) {
|
||||
t.Helper()
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
gotName, _ := first["name"].(string)
|
||||
if strings.TrimSpace(gotName) != expectedName {
|
||||
t.Fatalf("expected failure name %q, got %q", expectedName, gotName)
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, reasonSnippet) {
|
||||
t.Fatalf("expected failure reason to contain %q, got %q", reasonSnippet, reason)
|
||||
}
|
||||
}
|
||||
|
||||
func seedForwardForBatchAction(t *testing.T, repo *repo.Repository, opts batchForwardSeedOptions) int64 {
|
||||
t.Helper()
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'batch_action_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, opts.Now, opts.Now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, opts.TunnelName, 1.0, 1, "tls", 99999, opts.Now, opts.Now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, opts.TunnelName)
|
||||
|
||||
if opts.CreateUserTunnel {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(20, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, ?)
|
||||
`, tunnelID, opts.UserTunnelStatus).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(2, 'batch_action_user', ?, ?, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, opts.ForwardName, tunnelID, opts.Now, opts.Now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
return mustLastInsertID(t, repo, opts.ForwardName)
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestForwardBatchRedeployReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'batch_redeploy_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "batch-redeploy-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "batch-redeploy-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, 0)
|
||||
`, 2, "batch_redeploy_user", "redeploy-forward", tunnelID, "1.1.1.1:443", "fifo", now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, repo, "redeploy-forward")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/batch-redeploy", bytes.NewBufferString(`{"ids":[`+jsonNumber(forwardID)+`]}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected API success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected map result, got %T", out.Data)
|
||||
}
|
||||
if int(result["failCount"].(float64)) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %v", result["failCount"])
|
||||
}
|
||||
if int(result["successCount"].(float64)) != 0 {
|
||||
t.Fatalf("expected successCount=0, got %v", result["successCount"])
|
||||
}
|
||||
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
if gotName := strings.TrimSpace(first["name"].(string)); gotName != "redeploy-forward" {
|
||||
t.Fatalf("expected failure name redeploy-forward, got %q", gotName)
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, "转发入口端口不存在") {
|
||||
t.Fatalf("expected forward failure reason to mention missing entry port, got %q", reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelBatchRedeployReturnsFailureReasonsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "broken-redeploy-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "broken-redeploy-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "entry-only-node", "entry-only-secret", "10.0.0.20", "10.0.0.20", "", "20000-20010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
entryNodeID := mustLastInsertID(t, repo, "entry-only-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 20001, 'round', 1, 'tls')
|
||||
`, tunnelID, entryNodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/batch-redeploy", bytes.NewBufferString(`{"ids":[`+jsonNumber(tunnelID)+`]}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected API success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected map result, got %T", out.Data)
|
||||
}
|
||||
if int(result["failCount"].(float64)) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %v", result["failCount"])
|
||||
}
|
||||
|
||||
failures, ok := result["failures"].([]interface{})
|
||||
if !ok || len(failures) != 1 {
|
||||
t.Fatalf("expected exactly one failure detail, got %#v", result["failures"])
|
||||
}
|
||||
first, ok := failures[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected failure detail object, got %T", failures[0])
|
||||
}
|
||||
if gotName := strings.TrimSpace(first["name"].(string)); gotName != "broken-redeploy-tunnel" {
|
||||
t.Fatalf("expected failure name broken-redeploy-tunnel, got %q", gotName)
|
||||
}
|
||||
reason, _ := first["reason"].(string)
|
||||
if !strings.Contains(reason, "转发链目标不能为空") {
|
||||
t.Fatalf("expected tunnel failure reason to mention missing target, got %q", reason)
|
||||
}
|
||||
}
|
||||
@@ -1,19 +0,0 @@
|
||||
package contract
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func mustLastInsertID(t *testing.T, r *repo.Repository, label string) int64 {
|
||||
t.Helper()
|
||||
var id int64
|
||||
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
|
||||
t.Fatalf("read last_insert_rowid for %s: %v", label, err)
|
||||
}
|
||||
if id <= 0 {
|
||||
t.Fatalf("invalid last_insert_rowid for %s: %d", label, id)
|
||||
}
|
||||
return id
|
||||
}
|
||||
@@ -2,6 +2,8 @@ package contract_test
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
@@ -117,3 +119,43 @@ func tryQueryInt(t *testing.T, r *repo.Repository, query string, args ...interfa
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func valueAsInt(v interface{}) int {
|
||||
switch n := v.(type) {
|
||||
case float64:
|
||||
return int(n)
|
||||
case int:
|
||||
return n
|
||||
case int64:
|
||||
return int(n)
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func valueAsString(v interface{}) string {
|
||||
s, _ := v.(string)
|
||||
return s
|
||||
}
|
||||
|
||||
func valueAsBool(v interface{}) bool {
|
||||
switch b := v.(type) {
|
||||
case bool:
|
||||
return b
|
||||
case float64:
|
||||
return b != 0
|
||||
case int:
|
||||
return b != 0
|
||||
case int64:
|
||||
return b != 0
|
||||
case string:
|
||||
s := strings.TrimSpace(strings.ToLower(b))
|
||||
return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y"
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func jsonInt64(v int64) string {
|
||||
return strconv.FormatInt(v, 10)
|
||||
}
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
package contract
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
@@ -13,15 +13,12 @@ import (
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
httpserver "go-backend/internal/http"
|
||||
"go-backend/internal/http/handler"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestDiagnosisChainCoverageContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupDiagnosisContractRouter(t, secret)
|
||||
router, r := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
@@ -193,9 +190,129 @@ func TestDiagnosisChainCoverageContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestForwardDiagnosisRespectsTunnelIPPreferenceContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'normal_user', '3c85cdebade1c51cf64ca9f3c09d182d', 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(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, ip_preference)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "diagnose-ip-pref-forward", 1.0, 2, "tls", 99999, now, now, 1, nil, 0, "v6").Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "diagnose-ip-pref-forward")
|
||||
|
||||
insertNode := func(name, v4, v6 string) int64 {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, name, name+"-secret", v4, v4, v6, "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
return mustLastInsertID(t, r, name)
|
||||
}
|
||||
|
||||
entryNodeID := insertNode("entry-node-v6", "10.10.1.10", "2001:db8:10::10")
|
||||
chainNodeID := insertNode("chain-node-v6", "10.10.1.20", "2001:db8:10::20")
|
||||
exitNodeID := insertNode("exit-node-v6", "10.10.1.30", "2001:db8:10::30")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 30001, 'round', 1, 'tls')
|
||||
`, tunnelID, entryNodeID).Error; err != nil {
|
||||
t.Fatalf("insert entry chain: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 2, ?, 30002, 'round', 1, 'tls')
|
||||
`, tunnelID, chainNodeID).Error; err != nil {
|
||||
t.Fatalf("insert middle chain: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 3, ?, 30003, 'round', 1, 'tls')
|
||||
`, tunnelID, exitNodeID).Error; err != nil {
|
||||
t.Fatalf("insert exit chain: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||
`, 2, "normal_user", "ip-pref-forward", tunnelID, "8.8.8.8:53", "fifo", now, now, 0).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, "ip-pref-forward")
|
||||
|
||||
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate user token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/diagnose", bytes.NewBufferString(`{"forwardId":`+strconv.FormatInt(forwardID, 10)+`}`))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
payload, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected object payload, got %T", out.Data)
|
||||
}
|
||||
results, ok := payload["results"].([]interface{})
|
||||
if !ok || len(results) == 0 {
|
||||
t.Fatalf("expected non-empty results, got %v", payload["results"])
|
||||
}
|
||||
|
||||
hasEntryToChain := false
|
||||
hasChainToExit := false
|
||||
for _, raw := range results {
|
||||
item, ok := raw.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
from := valueAsInt(item["fromChainType"])
|
||||
to := valueAsInt(item["toChainType"])
|
||||
targetIP := strings.TrimSpace(valueAsString(item["targetIp"]))
|
||||
|
||||
if from == 1 && to == 2 {
|
||||
hasEntryToChain = true
|
||||
if targetIP != "2001:db8:10::20" {
|
||||
t.Fatalf("expected entry->chain diagnosis target to use IPv6, got %q", targetIP)
|
||||
}
|
||||
}
|
||||
|
||||
if from == 2 && to == 3 {
|
||||
hasChainToExit = true
|
||||
if targetIP != "2001:db8:10::30" {
|
||||
t.Fatalf("expected chain->exit diagnosis target to use IPv6, got %q", targetIP)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !hasEntryToChain || !hasChainToExit {
|
||||
t.Fatalf("expected entry->chain and chain->exit steps, got entry=%v chain=%v", hasEntryToChain, hasChainToExit)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupDiagnosisContractRouter(t, secret)
|
||||
router, r := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
remoteToken := "remote-diagnose-token"
|
||||
@@ -346,53 +463,166 @@ func TestDiagnosisUsesFederationRuntimeForRemoteNodes(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func valueAsInt(v interface{}) int {
|
||||
switch n := v.(type) {
|
||||
case float64:
|
||||
return int(n)
|
||||
case int:
|
||||
return n
|
||||
case int64:
|
||||
return int(n)
|
||||
default:
|
||||
return 0
|
||||
func TestTunnelDiagnosisUsesConfiguredConnectIPContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
insertNode := func(name, ip string) int64 {
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, name, name+"-secret", ip, ip, "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
return mustLastInsertID(t, r, name)
|
||||
}
|
||||
}
|
||||
|
||||
func valueAsString(v interface{}) string {
|
||||
s, _ := v.(string)
|
||||
return s
|
||||
}
|
||||
entryNodeID := insertNode("entry-connectip", "10.80.0.10")
|
||||
middleNodeID := insertNode("middle-connectip", "10.80.0.20")
|
||||
exitNodeID := insertNode("exit-connectip", "10.80.0.30")
|
||||
|
||||
func valueAsBool(v interface{}) bool {
|
||||
switch b := v.(type) {
|
||||
case bool:
|
||||
return b
|
||||
case float64:
|
||||
return b != 0
|
||||
case int:
|
||||
return b != 0
|
||||
case int64:
|
||||
return b != 0
|
||||
case string:
|
||||
s := strings.TrimSpace(strings.ToLower(b))
|
||||
return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y"
|
||||
default:
|
||||
return false
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "diagnose-connectip-tunnel", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "diagnose-connectip-tunnel")
|
||||
|
||||
func setupDiagnosisContractRouter(t *testing.T, jwtSecret string) (http.Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
dbPath := filepath.Join(t.TempDir(), "diagnosis-contract.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 30001, 'round', 1, 'tls')
|
||||
`, tunnelID, entryNodeID).Error; err != nil {
|
||||
t.Fatalf("insert entry chain: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol, connect_ip)
|
||||
VALUES(?, 2, ?, 30002, 'round', 1, 'tls', ?)
|
||||
`, tunnelID, middleNodeID, "10.99.0.22").Error; err != nil {
|
||||
t.Fatalf("insert middle chain: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol, connect_ip)
|
||||
VALUES(?, 3, ?, 30003, 'round', 1, 'tls', ?)
|
||||
`, tunnelID, exitNodeID, "10.99.0.33").Error; err != nil {
|
||||
t.Fatalf("insert exit chain: %v", err)
|
||||
}
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
|
||||
t.Run("normal diagnose should use configured connectIp", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose", bytes.NewBufferString(`{"tunnelId":`+strconv.FormatInt(tunnelID, 10)+`}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
payload, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected object payload, got %T", out.Data)
|
||||
}
|
||||
results, ok := payload["results"].([]interface{})
|
||||
if !ok || len(results) == 0 {
|
||||
t.Fatalf("expected non-empty results, got %v", payload["results"])
|
||||
}
|
||||
|
||||
entryToMiddleOK := false
|
||||
middleToExitOK := false
|
||||
for _, raw := range results {
|
||||
item, ok := raw.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
from := valueAsInt(item["fromChainType"])
|
||||
to := valueAsInt(item["toChainType"])
|
||||
targetIP := strings.TrimSpace(valueAsString(item["targetIp"]))
|
||||
|
||||
if from == 1 && to == 2 && targetIP == "10.99.0.22" {
|
||||
entryToMiddleOK = true
|
||||
}
|
||||
if from == 2 && to == 3 && targetIP == "10.99.0.33" {
|
||||
middleToExitOK = true
|
||||
}
|
||||
}
|
||||
|
||||
if !entryToMiddleOK || !middleToExitOK {
|
||||
t.Fatalf("expected connectIp targets 10.99.0.22/10.99.0.33, got entry=%v middle=%v", entryToMiddleOK, middleToExitOK)
|
||||
}
|
||||
})
|
||||
|
||||
h := handler.New(r, jwtSecret)
|
||||
return httpserver.NewRouter(h, jwtSecret), r
|
||||
t.Run("stream diagnose start items should use configured connectIp", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/diagnose/stream", bytes.NewBufferString(`{"tunnelId":`+strconv.FormatInt(tunnelID, 10)+`}`))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", res.Code)
|
||||
}
|
||||
|
||||
scanner := bufio.NewScanner(bytes.NewReader(res.Body.Bytes()))
|
||||
startFound := false
|
||||
entryToMiddleOK := false
|
||||
middleToExitOK := false
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
var event map[string]interface{}
|
||||
if err := json.Unmarshal([]byte(line), &event); err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(valueAsString(event["type"])) != "start" {
|
||||
continue
|
||||
}
|
||||
startFound = true
|
||||
data, ok := event["data"].(map[string]interface{})
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
items, ok := data["items"].([]interface{})
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
for _, raw := range items {
|
||||
item, ok := raw.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
from := valueAsInt(item["fromChainType"])
|
||||
to := valueAsInt(item["toChainType"])
|
||||
targetIP := strings.TrimSpace(valueAsString(item["targetIp"]))
|
||||
if from == 1 && to == 2 && targetIP == "10.99.0.22" {
|
||||
entryToMiddleOK = true
|
||||
}
|
||||
if from == 2 && to == 3 && targetIP == "10.99.0.33" {
|
||||
middleToExitOK = true
|
||||
}
|
||||
}
|
||||
break
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
t.Fatalf("scan stream body: %v", err)
|
||||
}
|
||||
if !startFound {
|
||||
t.Fatalf("expected start event in stream response")
|
||||
}
|
||||
if !entryToMiddleOK || !middleToExitOK {
|
||||
t.Fatalf("expected start items with connectIp targets 10.99.0.22/10.99.0.33, got entry=%v middle=%v", entryToMiddleOK, middleToExitOK)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -166,7 +166,7 @@ func TestFederationDualPanelMiddleExitAutoPortContract(t *testing.T) {
|
||||
|
||||
assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, middleShareID, 1)
|
||||
assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, exitShareID, 1)
|
||||
assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ?`, entryShareID, 0)
|
||||
assertCount(t, providerRepo, `SELECT COUNT(1) FROM peer_share_runtime WHERE share_id = ? AND status = 1 AND applied = 1`, entryShareID, 1)
|
||||
}
|
||||
|
||||
func TestFederationDualPanelRemoteDiagnosisContract(t *testing.T) {
|
||||
@@ -624,42 +624,6 @@ func waitNodeStatus(t *testing.T, r *repo.Repository, nodeID int64, expectedStat
|
||||
}
|
||||
}
|
||||
|
||||
func valueAsInt(v interface{}) int {
|
||||
switch n := v.(type) {
|
||||
case float64:
|
||||
return int(n)
|
||||
case int:
|
||||
return n
|
||||
case int64:
|
||||
return int(n)
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func valueAsString(v interface{}) string {
|
||||
s, _ := v.(string)
|
||||
return s
|
||||
}
|
||||
|
||||
func valueAsBool(v interface{}) bool {
|
||||
switch b := v.(type) {
|
||||
case bool:
|
||||
return b
|
||||
case float64:
|
||||
return b != 0
|
||||
case int:
|
||||
return b != 0
|
||||
case int64:
|
||||
return b != 0
|
||||
case string:
|
||||
s := strings.TrimSpace(strings.ToLower(b))
|
||||
return s == "1" || s == "t" || s == "true" || s == "yes" || s == "y"
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) {
|
||||
providerSecret := "provider-portrange-jwt"
|
||||
providerRouter, providerRepo := setupContractRouter(t, providerSecret)
|
||||
@@ -759,6 +723,21 @@ func TestFederationRuntimeCommandPortRangeEnforcement(t *testing.T) {
|
||||
}
|
||||
|
||||
// Test: Non-service commands should pass through without port validation
|
||||
res = sendCommand("share-portrange-token", "UpdateLimiters", map[string]interface{}{
|
||||
"limiter": "federation-limit-test",
|
||||
"data": map[string]interface{}{
|
||||
"name": "federation-limit-test",
|
||||
"limits": []string{"$ 1MB 1MB"},
|
||||
},
|
||||
})
|
||||
out = response.R{}
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0 for UpdateLimiters command, got %d (msg: %s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
res = sendCommand("share-portrange-token", "reload", nil)
|
||||
out = response.R{}
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
|
||||
@@ -0,0 +1,693 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestFederationForwardCardFlowLinkageContract(t *testing.T) {
|
||||
secret := "federation-forward-flow-contract-jwt"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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, is_remote, remote_url, remote_token, remote_config)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "flow-local-node", "flow-local-secret", "10.20.30.40", "10.20.30.40", "", "32000-32020", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil {
|
||||
t.Fatalf("insert local node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "flow-local-node")
|
||||
|
||||
shareToken := "flow-linkage-share-token"
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "flow-linkage-share",
|
||||
NodeID: nodeID,
|
||||
Token: shareToken,
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 1536,
|
||||
ExpiryTime: 0,
|
||||
PortRangeStart: 32000,
|
||||
PortRangeEnd: 32020,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken(shareToken)
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
tunnelName := fmt.Sprintf("Share-%d-Port-%d", share.ID, 32001)
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, tunnelName, 1, "tcp", 1, now, now, 1, "", 0).Error; err != nil {
|
||||
t.Fatalf("insert share tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "flow-share-tunnel")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "admin_user", "flow-linkage-forward", tunnelID, "1.1.1.1:443", "fifo", 0, 0, now, now, 1, 0).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, "flow-linkage-forward")
|
||||
|
||||
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, 32001).Error; err != nil {
|
||||
t.Fatalf("insert forward_port: %v", err)
|
||||
}
|
||||
|
||||
forwardOut := requestContractEnvelope(t, router, adminToken, "/api/v1/forward/list", nil)
|
||||
if forwardOut.Code != 0 {
|
||||
t.Fatalf("forward list failed: code=%d msg=%q", forwardOut.Code, forwardOut.Msg)
|
||||
}
|
||||
forwardRows := mustContractSlice(t, forwardOut.Data, "forward list data")
|
||||
|
||||
var targetForward map[string]interface{}
|
||||
for _, row := range forwardRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if contractValueAsInt64(m["id"]) == forwardID {
|
||||
targetForward = m
|
||||
break
|
||||
}
|
||||
}
|
||||
if targetForward == nil {
|
||||
t.Fatalf("target forward %d not found in /forward/list response", forwardID)
|
||||
}
|
||||
|
||||
shareOut := requestContractEnvelope(t, router, adminToken, "/api/v1/federation/share/list", nil)
|
||||
if shareOut.Code != 0 {
|
||||
t.Fatalf("share list failed: code=%d msg=%q", shareOut.Code, shareOut.Msg)
|
||||
}
|
||||
localShareRows := mustContractSlice(t, shareOut.Data, "share list data")
|
||||
|
||||
remoteUsageOut := requestContractEnvelope(t, router, adminToken, "/api/v1/federation/share/remote-usage/list", nil)
|
||||
if remoteUsageOut.Code != 0 {
|
||||
t.Fatalf("remote usage list failed: code=%d msg=%q", remoteUsageOut.Code, remoteUsageOut.Msg)
|
||||
}
|
||||
remoteUsageRows := mustContractSlice(t, remoteUsageOut.Data, "remote usage data")
|
||||
if len(remoteUsageRows) != 0 {
|
||||
t.Fatalf("expected no remote usage rows in local-only fixture, got %d", len(remoteUsageRows))
|
||||
}
|
||||
|
||||
flowByShare := make(map[int64]int64)
|
||||
for _, row := range remoteUsageRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
shareID := contractValueAsInt64(m["shareId"])
|
||||
currentFlow := contractValueAsInt64(m["currentFlow"])
|
||||
if shareID > 0 && currentFlow > 0 {
|
||||
if currentFlow > flowByShare[shareID] {
|
||||
flowByShare[shareID] = currentFlow
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, row := range localShareRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
shareID := contractValueAsInt64(m["id"])
|
||||
currentFlow := contractValueAsInt64(m["currentFlow"])
|
||||
if shareID > 0 && currentFlow > 0 {
|
||||
if currentFlow > flowByShare[shareID] {
|
||||
flowByShare[shareID] = currentFlow
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
parsedShareID := contractParseShareIDFromTunnelName(contractValueAsString(targetForward["tunnelName"]))
|
||||
if parsedShareID != share.ID {
|
||||
t.Fatalf("expected parsed shareID=%d, got %d (tunnelName=%q)", share.ID, parsedShareID, contractValueAsString(targetForward["tunnelName"]))
|
||||
}
|
||||
|
||||
forwardCountByShare := make(map[int64]int)
|
||||
for _, row := range forwardRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
sid := contractParseShareIDFromTunnelName(contractValueAsString(m["tunnelName"]))
|
||||
if sid > 0 && flowByShare[sid] > 0 {
|
||||
forwardCountByShare[sid] = forwardCountByShare[sid] + 1
|
||||
}
|
||||
}
|
||||
|
||||
directFlow := contractValueAsInt64(targetForward["inFlow"]) + contractValueAsInt64(targetForward["outFlow"])
|
||||
if directFlow != 0 {
|
||||
t.Fatalf("fixture expectation failed: directFlow should be 0, got %d", directFlow)
|
||||
}
|
||||
|
||||
shareFlow := flowByShare[parsedShareID]
|
||||
if shareFlow <= 0 {
|
||||
t.Fatalf("expected merged share flow > 0 for share %d", parsedShareID)
|
||||
}
|
||||
|
||||
count := forwardCountByShare[parsedShareID]
|
||||
if count <= 0 {
|
||||
count = 1
|
||||
}
|
||||
estimated := shareFlow / int64(count)
|
||||
if estimated < 1 {
|
||||
estimated = 1
|
||||
}
|
||||
displayFlow := estimated
|
||||
|
||||
if displayFlow <= 0 {
|
||||
t.Fatalf("expected displayFlow > 0 after frontend-style merge, got %d", displayFlow)
|
||||
}
|
||||
if displayFlow != share.CurrentFlow {
|
||||
t.Fatalf("expected displayFlow=%d, got %d", share.CurrentFlow, displayFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationForwardCardFlowLinkageContractSplitShareFlowAcrossMultipleForwards(t *testing.T) {
|
||||
secret := "federation-forward-split-flow-contract-jwt"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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, is_remote, remote_url, remote_token, remote_config)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "flow-split-local-node", "flow-split-local-secret", "10.21.31.41", "10.21.31.41", "", "32100-32120", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 0, "", "", "").Error; err != nil {
|
||||
t.Fatalf("insert local node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "flow-split-local-node")
|
||||
|
||||
shareToken := "flow-split-share-token"
|
||||
if err := r.CreatePeerShare(&repo.PeerShare{
|
||||
Name: "flow-split-share",
|
||||
NodeID: nodeID,
|
||||
Token: shareToken,
|
||||
MaxBandwidth: 0,
|
||||
CurrentFlow: 4097,
|
||||
ExpiryTime: 0,
|
||||
PortRangeStart: 32100,
|
||||
PortRangeEnd: 32120,
|
||||
IsActive: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}); err != nil {
|
||||
t.Fatalf("create share: %v", err)
|
||||
}
|
||||
share, err := r.GetPeerShareByToken(shareToken)
|
||||
if err != nil || share == nil {
|
||||
t.Fatalf("load share: %v", err)
|
||||
}
|
||||
|
||||
createShareForward := func(name string, port int) int64 {
|
||||
t.Helper()
|
||||
|
||||
tunnelName := fmt.Sprintf("Share-%d-Port-%d", share.ID, port)
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, tunnelName, 1, "tcp", 1, now, now, 1, "", 0).Error; err != nil {
|
||||
t.Fatalf("insert share tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "flow-split-tunnel")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "admin_user", name, tunnelID, "1.1.1.1:443", "fifo", 0, 0, now, now, 1, 0).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, name)
|
||||
|
||||
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port).Error; err != nil {
|
||||
t.Fatalf("insert forward_port: %v", err)
|
||||
}
|
||||
|
||||
return forwardID
|
||||
}
|
||||
|
||||
forwardIDA := createShareForward("flow-split-forward-a", 32101)
|
||||
forwardIDB := createShareForward("flow-split-forward-b", 32102)
|
||||
|
||||
forwardOut := requestContractEnvelope(t, router, adminToken, "/api/v1/forward/list", nil)
|
||||
if forwardOut.Code != 0 {
|
||||
t.Fatalf("forward list failed: code=%d msg=%q", forwardOut.Code, forwardOut.Msg)
|
||||
}
|
||||
forwardRows := mustContractSlice(t, forwardOut.Data, "forward list data")
|
||||
|
||||
shareOut := requestContractEnvelope(t, router, adminToken, "/api/v1/federation/share/list", nil)
|
||||
if shareOut.Code != 0 {
|
||||
t.Fatalf("share list failed: code=%d msg=%q", shareOut.Code, shareOut.Msg)
|
||||
}
|
||||
localShareRows := mustContractSlice(t, shareOut.Data, "share list data")
|
||||
|
||||
remoteUsageOut := requestContractEnvelope(t, router, adminToken, "/api/v1/federation/share/remote-usage/list", nil)
|
||||
if remoteUsageOut.Code != 0 {
|
||||
t.Fatalf("remote usage list failed: code=%d msg=%q", remoteUsageOut.Code, remoteUsageOut.Msg)
|
||||
}
|
||||
remoteUsageRows := mustContractSlice(t, remoteUsageOut.Data, "remote usage data")
|
||||
if len(remoteUsageRows) != 0 {
|
||||
t.Fatalf("expected no remote usage rows in local-only fixture, got %d", len(remoteUsageRows))
|
||||
}
|
||||
|
||||
flowByShare := make(map[int64]int64)
|
||||
for _, row := range remoteUsageRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
shareID := contractValueAsInt64(m["shareId"])
|
||||
currentFlow := contractValueAsInt64(m["currentFlow"])
|
||||
if shareID > 0 && currentFlow > 0 {
|
||||
if currentFlow > flowByShare[shareID] {
|
||||
flowByShare[shareID] = currentFlow
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, row := range localShareRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
shareID := contractValueAsInt64(m["id"])
|
||||
currentFlow := contractValueAsInt64(m["currentFlow"])
|
||||
if shareID > 0 && currentFlow > 0 {
|
||||
if currentFlow > flowByShare[shareID] {
|
||||
flowByShare[shareID] = currentFlow
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
shareFlow := flowByShare[share.ID]
|
||||
if shareFlow <= 0 {
|
||||
t.Fatalf("expected merged share flow > 0 for share %d", share.ID)
|
||||
}
|
||||
|
||||
forwardCountByShare := make(map[int64]int)
|
||||
for _, row := range forwardRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
sid := contractParseShareIDFromTunnelName(contractValueAsString(m["tunnelName"]))
|
||||
if sid > 0 && flowByShare[sid] > 0 {
|
||||
forwardCountByShare[sid] = forwardCountByShare[sid] + 1
|
||||
}
|
||||
}
|
||||
|
||||
count := forwardCountByShare[share.ID]
|
||||
if count != 2 {
|
||||
t.Fatalf("expected 2 forwards sharing share %d, got %d", share.ID, count)
|
||||
}
|
||||
|
||||
expectedEach := shareFlow / int64(count)
|
||||
if expectedEach < 1 {
|
||||
expectedEach = 1
|
||||
}
|
||||
|
||||
findForward := func(forwardID int64) map[string]interface{} {
|
||||
t.Helper()
|
||||
for _, row := range forwardRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if contractValueAsInt64(m["id"]) == forwardID {
|
||||
return m
|
||||
}
|
||||
}
|
||||
t.Fatalf("forward %d not found in /forward/list response", forwardID)
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, forwardID := range []int64{forwardIDA, forwardIDB} {
|
||||
forward := findForward(forwardID)
|
||||
sid := contractParseShareIDFromTunnelName(contractValueAsString(forward["tunnelName"]))
|
||||
if sid != share.ID {
|
||||
t.Fatalf("expected parsed shareID=%d, got %d for forward %d", share.ID, sid, forwardID)
|
||||
}
|
||||
|
||||
directFlow := contractValueAsInt64(forward["inFlow"]) + contractValueAsInt64(forward["outFlow"])
|
||||
if directFlow != 0 {
|
||||
t.Fatalf("fixture expectation failed: directFlow should be 0 for forward %d, got %d", forwardID, directFlow)
|
||||
}
|
||||
|
||||
displayFlow := int64(0)
|
||||
if directFlow > 0 {
|
||||
displayFlow = directFlow
|
||||
} else {
|
||||
shareFlowForForward := flowByShare[sid]
|
||||
if shareFlowForForward > 0 {
|
||||
cnt := forwardCountByShare[sid]
|
||||
if cnt <= 0 {
|
||||
cnt = 1
|
||||
}
|
||||
estimated := shareFlowForForward / int64(cnt)
|
||||
if estimated < 1 {
|
||||
estimated = 1
|
||||
}
|
||||
displayFlow = estimated
|
||||
}
|
||||
}
|
||||
|
||||
if displayFlow <= 0 {
|
||||
t.Fatalf("expected displayFlow > 0 for forward %d, got %d", forwardID, displayFlow)
|
||||
}
|
||||
if displayFlow != expectedEach {
|
||||
t.Fatalf("expected displayFlow=%d for forward %d, got %d", expectedEach, forwardID, displayFlow)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationForwardCardFlowLinkageContractResolvesShareByTunnelBindingWhenTunnelNameIsCustom(t *testing.T) {
|
||||
secret := "federation-forward-binding-flow-contract-jwt"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
remoteShareID := int64(901)
|
||||
remoteShareFlow := int64(5000)
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(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, is_remote, remote_url, remote_token, remote_config)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`,
|
||||
"flow-binding-remote-node", "flow-binding-remote-secret", "10.31.41.51", "10.31.41.51", "", "33000-33020", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0, 1, "", "", fmt.Sprintf(`{"shareId":%d,"maxBandwidth":0,"currentFlow":%d,"portRangeStart":33000,"portRangeEnd":33020}`, remoteShareID, remoteShareFlow),
|
||||
).Error; err != nil {
|
||||
t.Fatalf("insert remote node: %v", err)
|
||||
}
|
||||
remoteNodeID := mustLastInsertID(t, r, "flow-binding-remote-node")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "federation-port-forward-custom-name", 1, "tcp", 1, now, now, 1, "", 0).Error; err != nil {
|
||||
t.Fatalf("insert custom tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "flow-binding-custom-tunnel")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "admin_user", "flow-binding-forward", tunnelID, "1.1.1.1:443", "fifo", 0, 0, now, now, 1, 0).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, "flow-binding-forward")
|
||||
|
||||
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, remoteNodeID, 33001).Error; err != nil {
|
||||
t.Fatalf("insert forward_port: %v", err)
|
||||
}
|
||||
|
||||
forwardOut := requestContractEnvelope(t, router, adminToken, "/api/v1/forward/list", nil)
|
||||
if forwardOut.Code != 0 {
|
||||
t.Fatalf("forward list failed: code=%d msg=%q", forwardOut.Code, forwardOut.Msg)
|
||||
}
|
||||
forwardRows := mustContractSlice(t, forwardOut.Data, "forward list data")
|
||||
|
||||
shareOut := requestContractEnvelope(t, router, adminToken, "/api/v1/federation/share/list", nil)
|
||||
if shareOut.Code != 0 {
|
||||
t.Fatalf("share list failed: code=%d msg=%q", shareOut.Code, shareOut.Msg)
|
||||
}
|
||||
localShareRows := mustContractSlice(t, shareOut.Data, "share list data")
|
||||
|
||||
remoteUsageOut := requestContractEnvelope(t, router, adminToken, "/api/v1/federation/share/remote-usage/list", nil)
|
||||
if remoteUsageOut.Code != 0 {
|
||||
t.Fatalf("remote usage list failed: code=%d msg=%q", remoteUsageOut.Code, remoteUsageOut.Msg)
|
||||
}
|
||||
remoteUsageRows := mustContractSlice(t, remoteUsageOut.Data, "remote usage data")
|
||||
if len(remoteUsageRows) == 0 {
|
||||
t.Fatalf("expected non-empty remote usage rows")
|
||||
}
|
||||
|
||||
findForward := func(id int64) map[string]interface{} {
|
||||
t.Helper()
|
||||
for _, row := range forwardRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if contractValueAsInt64(m["id"]) == id {
|
||||
return m
|
||||
}
|
||||
}
|
||||
t.Fatalf("forward %d not found in /forward/list response", id)
|
||||
return nil
|
||||
}
|
||||
|
||||
flowByShare := make(map[int64]int64)
|
||||
shareIDsByTunnel := make(map[int64]map[int64]struct{})
|
||||
|
||||
for _, row := range remoteUsageRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
shareID := contractValueAsInt64(m["shareId"])
|
||||
currentFlow := contractValueAsInt64(m["currentFlow"])
|
||||
if shareID > 0 && currentFlow > 0 {
|
||||
if currentFlow > flowByShare[shareID] {
|
||||
flowByShare[shareID] = currentFlow
|
||||
}
|
||||
}
|
||||
|
||||
bindings, _ := m["bindings"].([]interface{})
|
||||
for _, bindingRaw := range bindings {
|
||||
binding, ok := bindingRaw.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
tunnelIDVal := contractValueAsInt64(binding["tunnelId"])
|
||||
chainType := contractValueAsInt64(binding["chainType"])
|
||||
if shareID <= 0 || tunnelIDVal <= 0 {
|
||||
continue
|
||||
}
|
||||
if chainType != 1 {
|
||||
continue
|
||||
}
|
||||
setByTunnel, ok := shareIDsByTunnel[tunnelIDVal]
|
||||
if !ok {
|
||||
setByTunnel = make(map[int64]struct{})
|
||||
shareIDsByTunnel[tunnelIDVal] = setByTunnel
|
||||
}
|
||||
setByTunnel[shareID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
for _, row := range localShareRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
shareID := contractValueAsInt64(m["id"])
|
||||
currentFlow := contractValueAsInt64(m["currentFlow"])
|
||||
if shareID > 0 && currentFlow > 0 {
|
||||
if currentFlow > flowByShare[shareID] {
|
||||
flowByShare[shareID] = currentFlow
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
targetForward := findForward(forwardID)
|
||||
parsedByName := contractParseShareIDFromTunnelName(contractValueAsString(targetForward["tunnelName"]))
|
||||
if parsedByName != 0 {
|
||||
t.Fatalf("expected custom tunnel name cannot be parsed as Share-*-Port-*, got %d", parsedByName)
|
||||
}
|
||||
|
||||
resolveShareIDForForward := func(forward map[string]interface{}) int64 {
|
||||
candidates := make(map[int64]struct{})
|
||||
|
||||
shareIDFromName := contractParseShareIDFromTunnelName(contractValueAsString(forward["tunnelName"]))
|
||||
if shareIDFromName > 0 {
|
||||
candidates[shareIDFromName] = struct{}{}
|
||||
}
|
||||
|
||||
tunnelIDVal := contractValueAsInt64(forward["tunnelId"])
|
||||
if setByTunnel, ok := shareIDsByTunnel[tunnelIDVal]; ok {
|
||||
for sid := range setByTunnel {
|
||||
candidates[sid] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
var bestShareID int64
|
||||
bestFlow := int64(0)
|
||||
for sid := range candidates {
|
||||
flow := flowByShare[sid]
|
||||
if flow > bestFlow {
|
||||
bestFlow = flow
|
||||
bestShareID = sid
|
||||
}
|
||||
}
|
||||
return bestShareID
|
||||
}
|
||||
|
||||
resolvedShareID := resolveShareIDForForward(targetForward)
|
||||
if resolvedShareID != remoteShareID {
|
||||
t.Fatalf("expected resolved shareID=%d via tunnel binding, got %d", remoteShareID, resolvedShareID)
|
||||
}
|
||||
|
||||
forwardCountByShare := make(map[int64]int)
|
||||
resolvedByForwardID := make(map[int64]int64)
|
||||
for _, row := range forwardRows {
|
||||
m, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
fid := contractValueAsInt64(m["id"])
|
||||
sid := resolveShareIDForForward(m)
|
||||
if sid > 0 {
|
||||
resolvedByForwardID[fid] = sid
|
||||
}
|
||||
if sid > 0 && flowByShare[sid] > 0 {
|
||||
forwardCountByShare[sid] = forwardCountByShare[sid] + 1
|
||||
}
|
||||
}
|
||||
|
||||
directFlow := contractValueAsInt64(targetForward["inFlow"]) + contractValueAsInt64(targetForward["outFlow"])
|
||||
if directFlow != 0 {
|
||||
t.Fatalf("fixture expectation failed: directFlow should be 0, got %d", directFlow)
|
||||
}
|
||||
|
||||
shareFlow := flowByShare[resolvedByForwardID[forwardID]]
|
||||
if shareFlow <= 0 {
|
||||
t.Fatalf("expected merged share flow > 0 for resolved share %d", resolvedByForwardID[forwardID])
|
||||
}
|
||||
|
||||
count := forwardCountByShare[resolvedByForwardID[forwardID]]
|
||||
if count <= 0 {
|
||||
count = 1
|
||||
}
|
||||
estimated := shareFlow / int64(count)
|
||||
if estimated < 1 {
|
||||
estimated = 1
|
||||
}
|
||||
|
||||
displayFlow := estimated
|
||||
if displayFlow <= 0 {
|
||||
t.Fatalf("expected displayFlow > 0 after tunnel-binding-based merge, got %d", displayFlow)
|
||||
}
|
||||
if displayFlow != remoteShareFlow {
|
||||
t.Fatalf("expected displayFlow=%d, got %d", remoteShareFlow, displayFlow)
|
||||
}
|
||||
}
|
||||
|
||||
func requestContractEnvelope(t *testing.T, router http.Handler, token string, path string, body interface{}) response.R {
|
||||
t.Helper()
|
||||
|
||||
payload := []byte("{}")
|
||||
if body != nil {
|
||||
raw, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request body for %s: %v", path, err)
|
||||
}
|
||||
payload = raw
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected http 200 for %s, got %d", path, res.Code)
|
||||
}
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response for %s: %v", path, err)
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
func mustContractSlice(t *testing.T, data interface{}, label string) []interface{} {
|
||||
t.Helper()
|
||||
|
||||
rows, ok := data.([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected %s to be []interface{}, got %T", label, data)
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func contractParseShareIDFromTunnelName(tunnelName string) int64 {
|
||||
normalized := strings.TrimSpace(tunnelName)
|
||||
if !strings.HasPrefix(normalized, "Share-") {
|
||||
return 0
|
||||
}
|
||||
raw := strings.TrimPrefix(normalized, "Share-")
|
||||
idx := strings.Index(raw, "-Port-")
|
||||
if idx <= 0 {
|
||||
return 0
|
||||
}
|
||||
shareID, err := strconv.ParseInt(strings.TrimSpace(raw[:idx]), 10, 64)
|
||||
if err != nil || shareID <= 0 {
|
||||
return 0
|
||||
}
|
||||
return shareID
|
||||
}
|
||||
|
||||
func contractValueAsInt64(v interface{}) int64 {
|
||||
switch n := v.(type) {
|
||||
case int64:
|
||||
return n
|
||||
case int:
|
||||
return int64(n)
|
||||
case float64:
|
||||
return int64(n)
|
||||
case json.Number:
|
||||
i, err := n.Int64()
|
||||
if err == nil {
|
||||
return i
|
||||
}
|
||||
f, err := n.Float64()
|
||||
if err == nil {
|
||||
return int64(f)
|
||||
}
|
||||
return 0
|
||||
case string:
|
||||
i, err := strconv.ParseInt(strings.TrimSpace(n), 10, 64)
|
||||
if err == nil {
|
||||
return i
|
||||
}
|
||||
return 0
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func contractValueAsString(v interface{}) string {
|
||||
s, _ := v.(string)
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
const contractBytesPerGB int64 = 1024 * 1024 * 1024
|
||||
|
||||
func TestForwardResumeBlockedWhenUserFlowExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(2)
|
||||
tunnelID := int64(1)
|
||||
forwardID := int64(1)
|
||||
|
||||
flowGB := int64(120)
|
||||
used := flowGB*contractBytesPerGB + 1
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'flow_user', 'pwd', 1, 2727251700000, ?, ?, 0, 1, 99999, ?, ?, 1)
|
||||
`, userID, flowGB, used, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'flow_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := repo.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(?, ?, 'flow_user', 'flow_forward', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 0, 0)
|
||||
`, forwardID, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "flow_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/resume", bytes.NewBufferString(`{"id":1}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when flow exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "流量") {
|
||||
t.Fatalf("expected flow exceeded message, got %q", out.Msg)
|
||||
}
|
||||
|
||||
status := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = ?`, forwardID)
|
||||
if status != 0 {
|
||||
t.Fatalf("expected forward status to remain 0, got %d", status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResumeBlockedWhenUserTunnelFlowExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(2)
|
||||
tunnelID := int64(1)
|
||||
forwardID := int64(1)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'ut_flow_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, userID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'ut_flow_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
utFlowGB := int64(120)
|
||||
utUsed := utFlowGB * contractBytesPerGB
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 99999, ?, ?, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID, utFlowGB, utUsed).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := repo.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(?, ?, 'ut_flow_user', 'ut_flow_forward', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 0, 0)
|
||||
`, forwardID, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "ut_flow_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/resume", bytes.NewBufferString(`{"id":1}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when tunnel flow exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "隧道") || !strings.Contains(out.Msg, "流量") {
|
||||
t.Fatalf("expected tunnel flow exceeded message, got %q", out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardCreateBlockedWhenFlowExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(2)
|
||||
tunnelID := int64(1)
|
||||
|
||||
flowGB := int64(120)
|
||||
used := flowGB*contractBytesPerGB + 1
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'create_flow_user', 'pwd', 1, 2727251700000, ?, ?, 0, 1, 99999, ?, ?, 1)
|
||||
`, userID, flowGB, used, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'create_flow_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "create_flow_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
payload := `{"tunnelId":1,"name":"n","remoteAddr":"1.1.1.1:53"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when flow exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "流量") {
|
||||
t.Fatalf("expected flow exceeded message, got %q", out.Msg)
|
||||
}
|
||||
}
|
||||
@@ -2,10 +2,13 @@ package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -28,7 +31,7 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "contract-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
`, "contract-tunnel", 2.5, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "contract-tunnel")
|
||||
@@ -108,9 +111,20 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
||||
if !ok {
|
||||
t.Fatalf("expected object item, got %T", arr[0])
|
||||
}
|
||||
if got := int64(item["id"].(float64)); got != userForwardID {
|
||||
idFloat, ok := item["id"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected id to be float64, got %T", item["id"])
|
||||
}
|
||||
if got := int64(idFloat); got != userForwardID {
|
||||
t.Fatalf("expected forward id %d, got %d", userForwardID, got)
|
||||
}
|
||||
ratioFloat, ok := item["tunnelTrafficRatio"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected tunnelTrafficRatio to be float64, got %T", item["tunnelTrafficRatio"])
|
||||
}
|
||||
if ratioFloat != 2.5 {
|
||||
t.Fatalf("expected tunnelTrafficRatio 2.5, got %v", ratioFloat)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("forward diagnose returns structured payload", func(t *testing.T) {
|
||||
@@ -143,7 +157,11 @@ func TestForwardOwnershipAndScopeContracts(t *testing.T) {
|
||||
if _, ok := first["message"]; !ok {
|
||||
t.Fatalf("expected message field in diagnosis result")
|
||||
}
|
||||
if got := int(first["fromChainType"].(float64)); got != 1 {
|
||||
fromChainTypeFloat, ok := first["fromChainType"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected fromChainType to be float64, got %T", first["fromChainType"])
|
||||
}
|
||||
if got := int(fromChainTypeFloat); got != 1 {
|
||||
t.Fatalf("expected fromChainType=1, got %d", got)
|
||||
}
|
||||
})
|
||||
@@ -471,6 +489,896 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserTunnelSaveIgnoresDeletedSpeedLimitContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(101, 'user_tunnel_speed_user_a', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "user-tunnel-missing-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "user-tunnel-missing-speed-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "user-tunnel-missing-speed-limit", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, repo, "user-tunnel-missing-speed-limit")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(31, 101, ?, ?, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID, speedID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, speedID).Error; err != nil {
|
||||
t.Fatalf("delete speed limit: %v", err)
|
||||
}
|
||||
|
||||
t.Run("user tunnel update auto clears missing speed", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": 31,
|
||||
"flow": 99999,
|
||||
"num": 999,
|
||||
"expTime": int64(2727251700000),
|
||||
"flowResetTime": 1,
|
||||
"status": 1,
|
||||
"speedId": speedID,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", adminToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
assertCode(t, updateRes, 0)
|
||||
|
||||
var updatedSpeed sql.NullInt64
|
||||
if err := repo.DB().Raw(`SELECT speed_id FROM user_tunnel WHERE id = 31`).Row().Scan(&updatedSpeed); err != nil {
|
||||
t.Fatalf("query updated user_tunnel speed_id: %v", err)
|
||||
}
|
||||
if updatedSpeed.Valid {
|
||||
t.Fatalf("expected updated user_tunnel speed_id to be NULL, got %d", updatedSpeed.Int64)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("user tunnel batch assign auto clears missing speed", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`UPDATE user_tunnel SET speed_id = ? WHERE id = 31`, speedID).Error; err != nil {
|
||||
t.Fatalf("prepare user_tunnel speed_id for batch assign: %v", err)
|
||||
}
|
||||
|
||||
assignPayload := map[string]interface{}{
|
||||
"userId": 101,
|
||||
"tunnels": []map[string]interface{}{{
|
||||
"tunnelId": tunnelID,
|
||||
"speedId": speedID,
|
||||
}},
|
||||
}
|
||||
assignBody, err := json.Marshal(assignPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal assign payload: %v", err)
|
||||
}
|
||||
assignReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", bytes.NewReader(assignBody))
|
||||
assignReq.Header.Set("Authorization", adminToken)
|
||||
assignReq.Header.Set("Content-Type", "application/json")
|
||||
assignRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(assignRes, assignReq)
|
||||
assertCode(t, assignRes, 0)
|
||||
|
||||
var assignedSpeed sql.NullInt64
|
||||
if err := repo.DB().Raw(`SELECT speed_id FROM user_tunnel WHERE id = 31`).Row().Scan(&assignedSpeed); err != nil {
|
||||
t.Fatalf("query assigned user_tunnel speed_id: %v", err)
|
||||
}
|
||||
if assignedSpeed.Valid {
|
||||
t.Fatalf("expected assigned user_tunnel speed_id to be NULL, got %d", assignedSpeed.Int64)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestForwardSpeedIDWriteAndClearContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'speed_user', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "forward-speed-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-speed-node", "forward-speed-secret", "10.30.0.1", "10.30.0.1", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, "forward-speed-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 31001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "forward-speed-limit-a", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit a: %v", err)
|
||||
}
|
||||
speedIDA := mustLastInsertID(t, repo, "forward-speed-limit-a")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "forward-speed-limit-b", 4096, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit b: %v", err)
|
||||
}
|
||||
speedIDB := mustLastInsertID(t, repo, "forward-speed-limit-b")
|
||||
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
stopNode := startMockNodeSession(t, server.URL, "forward-speed-secret")
|
||||
defer stopNode()
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "forward-speed-target",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedIDA,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "forward-speed-target")
|
||||
storedSpeed := repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var createdSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&createdSpeed); err != nil {
|
||||
t.Fatalf("query created forward speed_id: %v", err)
|
||||
}
|
||||
if !createdSpeed.Valid || createdSpeed.Int64 != speedIDA {
|
||||
t.Fatalf("expected created speed_id=%d, got valid=%v value=%d", speedIDA, createdSpeed.Valid, createdSpeed.Int64)
|
||||
}
|
||||
|
||||
updateToBPayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"speedId": speedIDB,
|
||||
}
|
||||
updateToBBody, err := json.Marshal(updateToBPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update-to-b payload: %v", err)
|
||||
}
|
||||
updateToBReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateToBBody))
|
||||
updateToBReq.Header.Set("Authorization", adminToken)
|
||||
updateToBReq.Header.Set("Content-Type", "application/json")
|
||||
updateToBRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateToBRes, updateToBReq)
|
||||
assertCode(t, updateToBRes, 0)
|
||||
|
||||
storedSpeed = repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var updatedSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&updatedSpeed); err != nil {
|
||||
t.Fatalf("query updated forward speed_id: %v", err)
|
||||
}
|
||||
if !updatedSpeed.Valid || updatedSpeed.Int64 != speedIDB {
|
||||
t.Fatalf("expected updated speed_id=%d, got valid=%v value=%d", speedIDB, updatedSpeed.Valid, updatedSpeed.Int64)
|
||||
}
|
||||
|
||||
clearPayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"speedId": nil,
|
||||
}
|
||||
clearBody, err := json.Marshal(clearPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal clear payload: %v", err)
|
||||
}
|
||||
clearReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(clearBody))
|
||||
clearReq.Header.Set("Authorization", adminToken)
|
||||
clearReq.Header.Set("Content-Type", "application/json")
|
||||
clearRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(clearRes, clearReq)
|
||||
assertCode(t, clearRes, 0)
|
||||
|
||||
storedSpeed = repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var clearedSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&clearedSpeed); err != nil {
|
||||
t.Fatalf("query cleared forward speed_id: %v", err)
|
||||
}
|
||||
if clearedSpeed.Valid {
|
||||
t.Fatalf("expected cleared speed_id to be NULL, got %d", clearedSpeed.Int64)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardUpdateIgnoresDeletedSpeedLimitContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-update-missing-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "forward-update-missing-speed-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-update-missing-speed-node", "forward-update-missing-speed-secret", "10.32.0.1", "10.32.0.1", "", "42000-42010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, "forward-update-missing-speed-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 42001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "forward-update-missing-speed-limit", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, repo, "forward-update-missing-speed-limit")
|
||||
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
stopNode := startMockNodeSession(t, server.URL, "forward-update-missing-speed-secret")
|
||||
defer stopNode()
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "forward-update-missing-speed-target",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedID,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "forward-update-missing-speed-target")
|
||||
|
||||
if err := repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, speedID).Error; err != nil {
|
||||
t.Fatalf("delete speed limit: %v", err)
|
||||
}
|
||||
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "forward-update-missing-speed-target-updated",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedID,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", adminToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
assertCode(t, updateRes, 0)
|
||||
|
||||
storedSpeed := repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var updatedSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&updatedSpeed); err != nil {
|
||||
t.Fatalf("query updated forward speed_id: %v", err)
|
||||
}
|
||||
if updatedSpeed.Valid {
|
||||
t.Fatalf("expected updated speed_id to be NULL after missing speed limit, got %d", updatedSpeed.Int64)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardCreateThenPauseResumeContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-toggle-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "forward-toggle-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-toggle-node", "forward-toggle-secret", "10.31.0.1", "10.31.0.1", "", "41000-41010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, "forward-toggle-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 41001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
stopNode := startMockNodeSession(t, server.URL, "forward-toggle-secret")
|
||||
defer stopNode()
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "forward-toggle-target",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "forward-toggle-target")
|
||||
|
||||
pauseBody, err := json.Marshal(map[string]interface{}{"id": forwardID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal pause payload: %v", err)
|
||||
}
|
||||
pauseReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/pause", bytes.NewReader(pauseBody))
|
||||
pauseReq.Header.Set("Authorization", adminToken)
|
||||
pauseReq.Header.Set("Content-Type", "application/json")
|
||||
pauseRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(pauseRes, pauseReq)
|
||||
assertCode(t, pauseRes, 0)
|
||||
|
||||
pausedStatus := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = ?`, forwardID)
|
||||
if pausedStatus != 0 {
|
||||
t.Fatalf("expected status=0 after pause, got %d", pausedStatus)
|
||||
}
|
||||
|
||||
resumeBody, err := json.Marshal(map[string]interface{}{"id": forwardID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal resume payload: %v", err)
|
||||
}
|
||||
resumeReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/resume", bytes.NewReader(resumeBody))
|
||||
resumeReq.Header.Set("Authorization", adminToken)
|
||||
resumeReq.Header.Set("Content-Type", "application/json")
|
||||
resumeRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(resumeRes, resumeReq)
|
||||
assertCode(t, resumeRes, 0)
|
||||
|
||||
resumedStatus := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = ?`, forwardID)
|
||||
if resumedStatus != 1 {
|
||||
t.Fatalf("expected status=1 after resume, got %d", resumedStatus)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardUpdateRecoversFromAddressInUseContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(202, 'forward_bind_retry_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-bind-retry-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "forward-bind-retry-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-bind-retry-node", "forward-bind-retry-secret", "10.42.0.1", "10.42.0.1", "", "44000-44010", "", "v1", 1, 1, 1, now, now, 1, "10.42.0.9", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, "forward-bind-retry-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 44001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(41, 202, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "forward-bind-retry-target",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
|
||||
var mu sync.Mutex
|
||||
counts := map[string]int{}
|
||||
var addServiceAddrs []string
|
||||
triggerConflict := false
|
||||
stopNode := startMockNodeSessionWithCommandRecorder(t, server.URL, "forward-bind-retry-secret", func(cmdType string, data json.RawMessage) (bool, string) {
|
||||
key := strings.ToLower(strings.TrimSpace(cmdType))
|
||||
mu.Lock()
|
||||
counts[key]++
|
||||
attempt := counts[key]
|
||||
if strings.EqualFold(strings.TrimSpace(cmdType), "AddService") || strings.EqualFold(strings.TrimSpace(cmdType), "UpdateService") {
|
||||
var services []map[string]interface{}
|
||||
if err := json.Unmarshal(data, &services); err == nil {
|
||||
for _, svc := range services {
|
||||
if addr, _ := svc["addr"].(string); strings.TrimSpace(addr) != "" {
|
||||
addServiceAddrs = append(addServiceAddrs, addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
shouldFail := false
|
||||
if triggerConflict {
|
||||
if strings.EqualFold(strings.TrimSpace(cmdType), "UpdateService") && attempt == 1 {
|
||||
shouldFail = true
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(cmdType), "AddService") && attempt == 1 {
|
||||
shouldFail = true
|
||||
}
|
||||
}
|
||||
mu.Unlock()
|
||||
if shouldFail {
|
||||
return true, "create service 57_7_7_tcp failed: listen tcp4 0.0.0.0:46222: bind: address alreadyin use"
|
||||
}
|
||||
return false, ""
|
||||
})
|
||||
defer stopNode()
|
||||
waitNodeStatus(t, repo, nodeID, 1)
|
||||
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
mu.Lock()
|
||||
counts = map[string]int{}
|
||||
addServiceAddrs = nil
|
||||
triggerConflict = true
|
||||
mu.Unlock()
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "forward-bind-retry-target")
|
||||
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "forward-bind-retry-target-updated",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "9.9.9.9:8443",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", adminToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
assertCode(t, updateRes, 0)
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
boundPort := mustQueryInt(t, repo, `SELECT port FROM forward_port WHERE forward_id = ? LIMIT 1`, forwardID)
|
||||
if counts["updateservice"] != 1 {
|
||||
t.Fatalf("expected one UpdateService attempt, got %d (%v)", counts["updateservice"], counts)
|
||||
}
|
||||
if counts["deleteservice"] == 0 {
|
||||
t.Fatalf("expected DeleteService cleanup after address-in-use (%v)", counts)
|
||||
}
|
||||
if counts["addservice"] < 2 {
|
||||
t.Fatalf("expected AddService retry path to run at least twice total, got %d (%v)", counts["addservice"], counts)
|
||||
}
|
||||
foundBindAddr := false
|
||||
for _, addr := range addServiceAddrs {
|
||||
if addr == "10.42.0.9:"+strconv.Itoa(boundPort) {
|
||||
foundBindAddr = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !foundBindAddr {
|
||||
t.Fatalf("expected forward runtime to keep node listen addr 10.42.0.9:%d, got %v", boundPort, addServiceAddrs)
|
||||
}
|
||||
|
||||
storedRemoteAddr := mustQueryString(t, repo, `SELECT remote_addr FROM forward WHERE id = ?`, forwardID)
|
||||
if storedRemoteAddr != "9.9.9.9:8443" {
|
||||
t.Fatalf("expected remote_addr update to persist, got %q", storedRemoteAddr)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonNumber(v int64) string {
|
||||
return strconv.FormatInt(v, 10)
|
||||
}
|
||||
|
||||
func TestNonAdminCannotSetSpeedIdOrPort(t *testing.T) {
|
||||
secret := "contract-jwt-secret-perm"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'normal_user_perm', '3c85cdebade1c51cf64ca9f3c09d182d', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "perm-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "perm-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "perm-node", "perm-secret", "10.0.0.20", "10.0.0.20", "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
entryNodeID := mustLastInsertID(t, repo, "perm-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 30001, 'round', 1, 'tls')
|
||||
`, tunnelID, entryNodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := 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(?, ?, NULL, 10, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, 2, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 1)
|
||||
`, "perm-speed-limit", 2048, tunnelID, "perm-tunnel", now, now).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, repo, "perm-speed-limit")
|
||||
|
||||
userToken, err := auth.GenerateToken(2, "normal_user_perm", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate user token: %v", err)
|
||||
}
|
||||
|
||||
stopNode := startMockNodeSession(t, server.URL, "perm-secret")
|
||||
defer stopNode()
|
||||
|
||||
t.Run("non-admin cannot set speedId on create", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "perm-forward-speed",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.2.3.4:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedID,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, -1, "普通用户无法设置限速规则")
|
||||
})
|
||||
|
||||
t.Run("non-admin cannot set inPort out of range on create", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "perm-forward-port-out",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.2.3.4:443",
|
||||
"strategy": "fifo",
|
||||
"inPort": 12345,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code >= 0 {
|
||||
t.Errorf("expected port out of range error, got code=%d msg=%s", out.Code, out.Msg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-admin can set inPort within range on create", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "perm-forward-port-in",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.2.3.4:443",
|
||||
"strategy": "fifo",
|
||||
"inPort": 30005,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("non-admin can create without speedId and inPort", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "perm-forward-ok",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.2.3.4:443",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "perm-forward-ok")
|
||||
|
||||
t.Run("non-admin cannot update speedId", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "perm-forward-updated",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "5.6.7.8:443",
|
||||
"speedId": speedID,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, -1, "普通用户无法修改限速规则")
|
||||
})
|
||||
|
||||
t.Run("non-admin cannot update inPort out of range", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "perm-forward-updated2",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "5.6.7.8:443",
|
||||
"inPort": 54321,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code >= 0 {
|
||||
t.Errorf("expected port out of range error, got code=%d msg=%s", out.Code, out.Msg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("non-admin can update inPort within range", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "perm-forward-updated3",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "5.6.7.8:443",
|
||||
"inPort": 30006,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("non-admin can update without speedId and inPort", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "perm-forward-updated-ok",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "9.10.11.12:443",
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("non-admin can update when request keeps existing speedId", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`UPDATE forward SET speed_id = ? WHERE id = ?`, speedID, forwardID).Error; err != nil {
|
||||
t.Fatalf("assign forward speed limit: %v", err)
|
||||
}
|
||||
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "perm-forward-keep-speed",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "9.10.11.12:443",
|
||||
"speedId": speedID,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("non-admin can create with speedId null and inPort 0", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "perm-forward-null-values",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.2.3.4:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": nil,
|
||||
"inPort": 0,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("non-admin can update with speedId null", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "perm-forward-null-speed",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "9.10.11.12:443",
|
||||
"speedId": nil,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,346 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestForwardCreateBlockedWhenUserNumLimitExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(100)
|
||||
tunnelID := int64(1)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'num_limit_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 2, ?, ?, 1)
|
||||
`, userID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'num_limit_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(1, ?, 'num_limit_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 1: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(2, ?, 'num_limit_user', 'existing_forward_2', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 2: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "num_limit_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
payload := `{"tunnelId":1,"name":"new_forward","remoteAddr":"1.1.1.1:53"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when num limit exceeded, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
if !strings.Contains(out.Msg, "转发数量已达上限") {
|
||||
t.Fatalf("expected forward count limit message, got %q", out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResumeBlockedWhenUserNumLimitExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(101)
|
||||
tunnelID := int64(1)
|
||||
pausedForwardID := int64(3)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'num_resume_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 2, ?, ?, 1)
|
||||
`, userID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'num_resume_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(1, ?, 'num_resume_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 1: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(2, ?, 'num_resume_user', 'existing_forward_2', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 2: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(?, ?, 'num_resume_user', 'paused_forward', ?, '1.1.1.1:53', 'fifo', 0, 0, ?, ?, 0, 0)
|
||||
`, pausedForwardID, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert paused forward: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "num_resume_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/resume", bytes.NewBufferString(`{"id":3}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when num limit exceeded, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
if !strings.Contains(out.Msg, "转发数量已达上限") {
|
||||
t.Fatalf("expected forward count limit message, got %q", out.Msg)
|
||||
}
|
||||
|
||||
status := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = ?`, pausedForwardID)
|
||||
if status != 0 {
|
||||
t.Fatalf("expected forward status to remain 0, got %d", status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardCreateBlockedWhenUserTunnelNumLimitExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(102)
|
||||
tunnelID := int64(1)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'ut_num_limit_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, userID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'ut_num_limit_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 1, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel with num=1: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(1, ?, 'ut_num_limit_user', 'existing_tunnel_forward', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "ut_num_limit_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
payload := `{"tunnelId":1,"name":"new_tunnel_forward","remoteAddr":"1.1.1.1:53"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when user_tunnel num limit exceeded, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
if !strings.Contains(out.Msg, "隧道转发数量已达上限") {
|
||||
t.Fatalf("expected tunnel forward count limit message, got %q", out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardCreateAllowedWhenBelowUserNumLimit(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(103)
|
||||
tunnelID := int64(1)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'num_ok_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 3, ?, ?, 1)
|
||||
`, userID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'num_ok_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(1, ?, 'num_ok_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 1: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "num_ok_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
payload := `{"tunnelId":1,"name":"new_forward_ok","remoteAddr":"1.1.1.1:53"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected success (code=0) when below num limit, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardCreateAllowedWhenNumZero(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
userID := int64(104)
|
||||
tunnelID := int64(1)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(?, 'num_zero_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 0, ?, ?, 1)
|
||||
`, userID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 'num_zero_tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, ?, ?, NULL, 0, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, userID, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel with num=0: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(1, ?, 'num_zero_user', 'existing_forward_1', ?, '8.8.8.8:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 1: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.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(2, ?, 'num_zero_user', 'existing_forward_2', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, userID, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert existing forward 2: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(userID, "num_zero_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
payload := `{"tunnelId":1,"name":"new_forward_zero","remoteAddr":"1.1.1.1:53"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(payload))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected success (code=0) when num=0 (unlimited), got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestIssue313_EntryPortCrossTunnelConflictContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
insertNode := func(name, ip, portRange string) int64 {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, name, name+"-secret", ip, ip, "", portRange, "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
return mustLastInsertID(t, repo, name)
|
||||
}
|
||||
|
||||
entryB1 := insertNode("issue313-entry-b1", "10.100.0.2", "2000-2010")
|
||||
entryB2 := insertNode("issue313-entry-b2", "10.100.0.3", "2000-2010")
|
||||
chainA := insertNode("issue313-chain-a", "10.100.0.4", "3000-3010")
|
||||
chainB := insertNode("issue313-chain-b", "10.100.0.5", "3000-3010")
|
||||
exitA := insertNode("issue313-exit-a", "10.100.0.6", "4000-4010")
|
||||
exitB := insertNode("issue313-exit-b", "10.100.0.7", "4000-4010")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue313-tunnel-a", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel a: %v", err)
|
||||
}
|
||||
tunnelAID := mustLastInsertID(t, repo, "issue313-tunnel-a")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 2000, 'round', 1, 'tls')
|
||||
`, tunnelAID, entryB2).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel entry a: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 2, ?, 3000, 'round', 1, 'tls')
|
||||
`, tunnelAID, chainA).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel chain a: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 3, ?, 4000, 'round', 1, 'tls')
|
||||
`, tunnelAID, exitA).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel exit a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue313-tunnel-b", 1.0, 2, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel b: %v", err)
|
||||
}
|
||||
tunnelBID := mustLastInsertID(t, repo, "issue313-tunnel-b")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 2000, 'round', 1, 'tls')
|
||||
`, tunnelBID, entryB1).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel entry b1: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 2, ?, 3000, 'round', 1, 'tls')
|
||||
`, tunnelBID, chainB).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel chain b: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 3, ?, 4000, 'round', 1, 'tls')
|
||||
`, tunnelBID, exitB).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel exit b: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(3131, 1, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelAID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel for tunnel a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(1, 'admin_user', 'issue313-forward-a', ?, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, tunnelAID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward a: %v", err)
|
||||
}
|
||||
forwardAID := mustLastInsertID(t, repo, "issue313-forward-a")
|
||||
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardAID, entryB2, 2000).Error; err != nil {
|
||||
t.Fatalf("insert forward_port a: %v", err)
|
||||
}
|
||||
|
||||
// Simulate legacy dirty data: tunnel A already occupies port 2000 on entryB2.
|
||||
// When tunnel B adds entryB2, the inherited forward port should conflict cross-tunnel.
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardAID, entryB2, 2000).Error; err != nil {
|
||||
t.Fatalf("insert forward_port a on entryB2: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(3132, 1, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelBID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel for tunnel b: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(1, 'admin_user', 'issue313-forward-b', ?, '2.2.2.2:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, tunnelBID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward b: %v", err)
|
||||
}
|
||||
forwardBID := mustLastInsertID(t, repo, "issue313-forward-b")
|
||||
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardBID, entryB1, 2000).Error; err != nil {
|
||||
t.Fatalf("insert forward_port b: %v", err)
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"id": tunnelBID,
|
||||
"name": "issue313-tunnel-b",
|
||||
"type": 2,
|
||||
"flow": 99999,
|
||||
"trafficRatio": 1.0,
|
||||
"status": 1,
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": entryB1, "protocol": "tls", "strategy": "round"},
|
||||
{"nodeId": entryB2, "protocol": "tls", "strategy": "round"},
|
||||
},
|
||||
"chainNodes": []interface{}{
|
||||
[]map[string]interface{}{{"nodeId": chainB, "protocol": "tls", "strategy": "round"}},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": exitB, "protocol": "tls", "strategy": "round"},
|
||||
},
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected update failure due to cross-tunnel port conflict, got success with code 0")
|
||||
}
|
||||
|
||||
msgBytes := []byte(out.Msg)
|
||||
if !bytes.Contains(msgBytes, []byte("端口")) && !bytes.Contains(msgBytes, []byte("占用")) {
|
||||
t.Fatalf("expected port conflict error message, got %q", out.Msg)
|
||||
}
|
||||
|
||||
countB2 := mustQueryInt(t, repo, `SELECT COUNT(1) FROM forward_port WHERE forward_id = ? AND node_id = ?`, forwardBID, entryB2)
|
||||
if countB2 > 0 {
|
||||
t.Fatalf("expected no forward_port record for entryB2, but found %d", countB2)
|
||||
}
|
||||
|
||||
chainCountB2 := mustQueryInt(t, repo, `SELECT COUNT(1) FROM chain_tunnel WHERE tunnel_id = ? AND node_id = ?`, tunnelBID, entryB2)
|
||||
if chainCountB2 > 0 {
|
||||
t.Fatalf("expected no chain_tunnel record for entryB2, but found %d", chainCountB2)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestIssue349_ForwardListFormatsIPv6EntryAddressesContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue349-tunnel", 1.0, 1, "tcp", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "issue349-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue349-entry-node-a", "entry-secret-a", "2001:db8::10", "", "2001:db8::10", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node a: %v", err)
|
||||
}
|
||||
nodeAID := mustLastInsertID(t, repo, "issue349-entry-node-a")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "issue349-entry-node-b", "entry-secret-b", "2001:db8::30", "", "2001:db8::30", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 1).Error; err != nil {
|
||||
t.Fatalf("insert node b: %v", err)
|
||||
}
|
||||
nodeBID := mustLastInsertID(t, repo, "issue349-entry-node-b")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||
`, 1, "admin_user", "issue349-forward", tunnelID, "1.1.1.1:443", "fifo", now, now, 0).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, repo, "issue349-forward")
|
||||
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeAID, 32001).Error; err != nil {
|
||||
t.Fatalf("insert forward_port a: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port, in_ip) VALUES(?, ?, ?, ?)`, forwardID, nodeBID, 32002, "2001:db8::20").Error; err != nil {
|
||||
t.Fatalf("insert forward_port b: %v", err)
|
||||
}
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/forward/list", nil)
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("forward list failed: code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
rows := mustContractSlice(t, out.Data, "forward list data")
|
||||
var target map[string]interface{}
|
||||
for _, row := range rows {
|
||||
item, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if contractValueAsInt64(item["id"]) == forwardID {
|
||||
target = item
|
||||
break
|
||||
}
|
||||
}
|
||||
if target == nil {
|
||||
t.Fatalf("target forward %d not found in /forward/list response", forwardID)
|
||||
}
|
||||
|
||||
if got := contractValueAsString(target["inIp"]); got != "[2001:db8::10]:32001,[2001:db8::20]:32002" {
|
||||
t.Fatalf("expected bracketed IPv6 entry list, got %q", got)
|
||||
}
|
||||
if got := contractValueAsInt64(target["inPort"]); got != 32001 {
|
||||
t.Fatalf("expected first entry port 32001, got %d", got)
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -25,6 +25,7 @@ import (
|
||||
func TestCaptchaVerifyLoginContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := setupContractRouter(t, secret)
|
||||
verifiedToken := ""
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
@@ -34,7 +35,7 @@ func TestCaptchaVerifyLoginContract(t *testing.T) {
|
||||
t.Fatalf("enable captcha: %v", err)
|
||||
}
|
||||
|
||||
t.Run("login denied without verified captcha token", func(t *testing.T) {
|
||||
t.Run("login allowed when cloudflare keys are missing", func(t *testing.T) {
|
||||
body := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
@@ -42,10 +43,10 @@ func TestCaptchaVerifyLoginContract(t *testing.T) {
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertCodeMsg(t, resp, -1, "验证码校验失败")
|
||||
assertCode(t, resp, 0)
|
||||
})
|
||||
|
||||
t.Run("captcha token is one-time and consumed by login", func(t *testing.T) {
|
||||
t.Run("captcha verify remains compatible without cloudflare secret", func(t *testing.T) {
|
||||
verifyReq := httptest.NewRequest(http.MethodPost, "/api/v1/captcha/verify", bytes.NewBufferString(`{"id":"captcha-token-1","data":"ok"}`))
|
||||
verifyReq.Header.Set("Content-Type", "application/json")
|
||||
verifyResp := httptest.NewRecorder()
|
||||
@@ -65,14 +66,60 @@ func TestCaptchaVerifyLoginContract(t *testing.T) {
|
||||
t.Fatalf("unexpected captcha verify payload: success=%v token=%q", verifyOut.Success, verifyOut.Data.ValidToken)
|
||||
}
|
||||
|
||||
loginBody := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":"captcha-token-1"}`)
|
||||
verifiedToken = verifyOut.Data.ValidToken
|
||||
})
|
||||
|
||||
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
|
||||
`, "cloudflare_site_key", "test-site-key", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("set cloudflare site key: %v", err)
|
||||
}
|
||||
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
|
||||
`, "cloudflare_secret_key", "test-secret-key", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("set cloudflare secret key: %v", err)
|
||||
}
|
||||
|
||||
t.Run("login denied without verified captcha token", func(t *testing.T) {
|
||||
body := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","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, "验证码校验失败")
|
||||
})
|
||||
|
||||
t.Run("whmcs api client bypasses captcha", func(t *testing.T) {
|
||||
body := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":""}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-FLVX-API-Client", "whmcs")
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertCode(t, resp, 0)
|
||||
})
|
||||
|
||||
t.Run("captcha token is one-time and consumed by login", func(t *testing.T) {
|
||||
if strings.TrimSpace(verifiedToken) == "" {
|
||||
t.Fatalf("expected verified token from compatibility captcha verify")
|
||||
}
|
||||
|
||||
loginBody := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":"` + verifiedToken + `"}`)
|
||||
loginReq := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", loginBody)
|
||||
loginReq.Header.Set("Content-Type", "application/json")
|
||||
loginResp := httptest.NewRecorder()
|
||||
router.ServeHTTP(loginResp, loginReq)
|
||||
assertCode(t, loginResp, 0)
|
||||
|
||||
replayBody := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":"captcha-token-1"}`)
|
||||
replayBody := bytes.NewBufferString(`{"username":"admin_user","password":"admin_user","captchaId":"` + verifiedToken + `"}`)
|
||||
replayReq := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", replayBody)
|
||||
replayReq.Header.Set("Content-Type", "application/json")
|
||||
replayResp := httptest.NewRecorder()
|
||||
@@ -162,39 +209,24 @@ func TestOpenAPISubStoreContracts(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestSpeedLimitTunnelsRouteAlias(t *testing.T) {
|
||||
func TestSpeedLimitTunnelsRouteRemoved(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
|
||||
t.Run("missing token blocked", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
|
||||
resp := httptest.NewRecorder()
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
assertCodeMsg(t, resp, 401, "未登录或token已过期")
|
||||
})
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
t.Run("admin token receives success envelope", func(t *testing.T) {
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/tunnels", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
resp := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
})
|
||||
if resp.Code != http.StatusNotFound {
|
||||
t.Fatalf("expected status 404 after route removal, got %d", resp.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
@@ -652,15 +684,107 @@ func TestOpenMigratesLegacyNodeDualStackColumns(t *testing.T) {
|
||||
|
||||
columns := readTableColumns(t, r.DB(), "node")
|
||||
|
||||
for _, required := range []string{"server_ip_v4", "server_ip_v6", "inx"} {
|
||||
for _, required := range []string{"server_ip_v4", "server_ip_v6", "inx", "extra_ips"} {
|
||||
if !columns[required] {
|
||||
t.Fatalf("expected node column %q to exist after migration", required)
|
||||
}
|
||||
}
|
||||
|
||||
tunnelColumns := readTableColumns(t, r.DB(), "tunnel")
|
||||
if !tunnelColumns["inx"] {
|
||||
t.Fatalf("expected tunnel column %q to exist after migration", "inx")
|
||||
for _, required := range []string{"inx", "ip_preference"} {
|
||||
if !tunnelColumns[required] {
|
||||
t.Fatalf("expected tunnel column %q to exist after migration", required)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenMigratesVeryLegacyNodeAndTunnelColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "legacy-1.x.db")
|
||||
legacyDB, err := sql.Open("sqlite", dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open legacy sqlite: %v", err)
|
||||
}
|
||||
|
||||
t.Cleanup(func() {
|
||||
_ = legacyDB.Close()
|
||||
})
|
||||
|
||||
if _, err := legacyDB.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS node (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
secret VARCHAR(100) NOT NULL,
|
||||
server_ip VARCHAR(100) NOT NULL,
|
||||
port TEXT NOT NULL,
|
||||
interface_name VARCHAR(200),
|
||||
version VARCHAR(100),
|
||||
http INTEGER NOT NULL DEFAULT 0,
|
||||
tls INTEGER NOT NULL DEFAULT 0,
|
||||
socks INTEGER NOT NULL DEFAULT 0,
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER,
|
||||
status INTEGER NOT NULL
|
||||
)
|
||||
`); err != nil {
|
||||
t.Fatalf("create very legacy node table: %v", err)
|
||||
}
|
||||
|
||||
if _, err := legacyDB.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS tunnel (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
traffic_ratio REAL NOT NULL DEFAULT 1.0,
|
||||
type INTEGER NOT NULL,
|
||||
protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
|
||||
flow INTEGER NOT NULL,
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER NOT NULL,
|
||||
status INTEGER NOT NULL,
|
||||
in_ip TEXT
|
||||
)
|
||||
`); err != nil {
|
||||
t.Fatalf("create very legacy tunnel table: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if _, err := legacyDB.Exec(`
|
||||
INSERT INTO node(name, secret, server_ip, port, interface_name, version, http, tls, socks, created_time, updated_time, status)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "legacy-node", "legacy-secret", "10.10.0.1", "10000-10010", "eth0", "v-old", 1, 1, 1, now, now, 1); err != nil {
|
||||
t.Fatalf("seed legacy node row: %v", err)
|
||||
}
|
||||
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open migrated sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = r.Close()
|
||||
})
|
||||
|
||||
columns := readTableColumns(t, r.DB(), "node")
|
||||
for _, required := range []string{
|
||||
"server_ip_v4",
|
||||
"server_ip_v6",
|
||||
"extra_ips",
|
||||
"tcp_listen_addr",
|
||||
"udp_listen_addr",
|
||||
"inx",
|
||||
"is_remote",
|
||||
"remote_url",
|
||||
"remote_token",
|
||||
"remote_config",
|
||||
} {
|
||||
if !columns[required] {
|
||||
t.Fatalf("expected node column %q to exist after migration", required)
|
||||
}
|
||||
}
|
||||
|
||||
tunnelColumns := readTableColumns(t, r.DB(), "tunnel")
|
||||
for _, required := range []string{"inx", "ip_preference"} {
|
||||
if !tunnelColumns[required] {
|
||||
t.Fatalf("expected tunnel column %q to exist after migration", required)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,313 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestSpeedLimitWithoutTunnelContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
t.Run("create speed limit", func(t *testing.T) {
|
||||
body := `{"name":"test-limit-no-tunnel","speed":100,"status":1}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("list does not expose tunnel binding fields", func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d", out.Code)
|
||||
}
|
||||
|
||||
data, ok := out.Data.([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected data to be array, got %T", out.Data)
|
||||
}
|
||||
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if m["name"] != "test-limit-no-tunnel" {
|
||||
continue
|
||||
}
|
||||
if tunnelID, exists := m["tunnelId"]; exists && tunnelID != nil {
|
||||
t.Fatalf("expected tunnelId to be absent or nil, got %v", tunnelID)
|
||||
}
|
||||
if tunnelName, exists := m["tunnelName"]; exists && tunnelName != nil && tunnelName != "" {
|
||||
t.Fatalf("expected tunnelName to be absent or empty, got %v", tunnelName)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
t.Fatal("speed limit 'test-limit-no-tunnel' not found in list")
|
||||
})
|
||||
}
|
||||
|
||||
func TestSpeedLimitCreateIgnoresTunnelBindingContract(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)
|
||||
}
|
||||
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-speed-limit-create-ignore-tunnel")
|
||||
|
||||
body := `{"name":"test-limit-ignore-tunnel","speed":200,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/create", bytes.NewBufferString(body))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
assertCode(t, res, 0)
|
||||
|
||||
req = httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
res = httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d", out.Code)
|
||||
}
|
||||
|
||||
data, ok := out.Data.([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected data to be array, got %T", out.Data)
|
||||
}
|
||||
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if m["name"] != "test-limit-ignore-tunnel" {
|
||||
continue
|
||||
}
|
||||
if tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil {
|
||||
t.Fatalf("expected tunnelId ignored and nil, got %v", tunnelIDVal)
|
||||
}
|
||||
if tunnelNameVal, exists := m["tunnelName"]; exists && tunnelNameVal != nil && tunnelNameVal != "" {
|
||||
t.Fatalf("expected tunnelName ignored and empty, got %v", tunnelNameVal)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
t.Fatal("speed limit 'test-limit-ignore-tunnel' not found in list")
|
||||
}
|
||||
|
||||
func TestSpeedLimitUpdateIgnoresTunnelBindingContract(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)
|
||||
}
|
||||
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-speed-limit-update-ignore-tunnel")
|
||||
speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update-ignore-tunnel")
|
||||
|
||||
body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update-ignore-tunnel","speed":256,"tunnelId":` + jsonInt(tunnelID) + `,"status":1}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/update", bytes.NewBufferString(body))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
assertCode(t, res, 0)
|
||||
|
||||
req = httptest.NewRequest(http.MethodPost, "/api/v1/speed-limit/list", nil)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
res = httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d", out.Code)
|
||||
}
|
||||
|
||||
data, ok := out.Data.([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected data to be array, got %T", out.Data)
|
||||
}
|
||||
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if m["name"] != "test-limit-update-ignore-tunnel" {
|
||||
continue
|
||||
}
|
||||
if tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil {
|
||||
t.Fatalf("expected tunnelId ignored and nil after update, got %v", tunnelIDVal)
|
||||
}
|
||||
if speedVal, ok := m["speed"].(float64); !ok || int(speedVal) != 256 {
|
||||
t.Fatalf("expected speed 256 after update, got %v", m["speed"])
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
t.Fatal("speed limit 'test-limit-update-ignore-tunnel' not found in list")
|
||||
}
|
||||
|
||||
func TestSpeedLimitDatabaseNullableFields(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-null.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
id, err := r.CreateSpeedLimit("db-test-limit", 100, 1, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateSpeedLimit failed: %v", err)
|
||||
}
|
||||
if id <= 0 {
|
||||
t.Fatalf("expected valid id, got %d", id)
|
||||
}
|
||||
|
||||
var tunnelID sql.NullInt64
|
||||
var tunnelName sql.NullString
|
||||
err = r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE id = ?", id).Row().Scan(&tunnelID, &tunnelName)
|
||||
if err != nil {
|
||||
t.Fatalf("query failed: %v", err)
|
||||
}
|
||||
if tunnelID.Valid {
|
||||
t.Fatalf("expected TunnelID to be NULL, got %d", tunnelID.Int64)
|
||||
}
|
||||
if tunnelName.Valid && tunnelName.String != "" {
|
||||
t.Fatalf("expected TunnelName to be NULL or empty, got %s", tunnelName.String)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpeedLimitUpdateClearsHistoricalBinding(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-update-clear.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "speed-limit-update-clear-tunnel")
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?)
|
||||
`, "speed-limit-update-clear", 300, tunnelID, "speed-limit-update-clear-tunnel", now, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit with tunnel binding: %v", err)
|
||||
}
|
||||
speedLimitID := mustLastInsertID(t, r, "speed-limit-update-clear")
|
||||
|
||||
err = r.UpdateSpeedLimit(speedLimitID, "speed-limit-update-clear", 512, 1, time.Now().UnixMilli())
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateSpeedLimit failed: %v", err)
|
||||
}
|
||||
|
||||
var dbTunnelID sql.NullInt64
|
||||
var dbTunnelName sql.NullString
|
||||
err = r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE id = ?", speedLimitID).Row().Scan(&dbTunnelID, &dbTunnelName)
|
||||
if err != nil {
|
||||
t.Fatalf("query updated speed limit failed: %v", err)
|
||||
}
|
||||
if dbTunnelID.Valid {
|
||||
t.Fatalf("expected tunnel_id cleared after update, got %d", dbTunnelID.Int64)
|
||||
}
|
||||
if dbTunnelName.Valid && dbTunnelName.String != "" {
|
||||
t.Fatalf("expected tunnel_name cleared after update, got %q", dbTunnelName.String)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpeedLimitGetSpeed(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-getspeed.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, 1, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("create speed limit: %v", err)
|
||||
}
|
||||
|
||||
t.Run("GetSpeedLimitSpeed returns correct speed", func(t *testing.T) {
|
||||
speed, err := r.GetSpeedLimitSpeed(speedLimitID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetSpeedLimitSpeed failed: %v", err)
|
||||
}
|
||||
if speed != 500 {
|
||||
t.Fatalf("expected speed 500, got %d", speed)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GetSpeedLimitSpeed returns error for non-existent id", func(t *testing.T) {
|
||||
_, err := r.GetSpeedLimitSpeed(99999)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for non-existent speed limit ID")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func mustCreateSpeedLimitTunnel(t *testing.T, r *repo.Repository, name string) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, name, now, now).Error; err != nil {
|
||||
t.Fatalf("create tunnel failed: %v", err)
|
||||
}
|
||||
return mustLastInsertID(t, r, name)
|
||||
}
|
||||
|
||||
func mustCreateSpeedLimitRepo(t *testing.T, r *repo.Repository, name string) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
id, err := r.CreateSpeedLimit(name, 100, now, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("create speed limit failed: %v", err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
storeRepo "go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestTunnelDeletePreviewIncludesDependentRulesContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
sourceTunnelID, sourceNodeID := seedTunnelDeleteTunnelWithNode(t, repo, now, "preview-source-tunnel", "preview-source-node", "21000-21010")
|
||||
seedTunnelDeleteForward(t, repo, now, sourceTunnelID, sourceNodeID, "preview-forward", 21001)
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/tunnel/delete-preview", map[string]interface{}{"id": sourceTunnelID})
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected success, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
data, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected preview data object, got %T", out.Data)
|
||||
}
|
||||
if contractValueAsInt64(data["tunnelId"]) != sourceTunnelID {
|
||||
t.Fatalf("unexpected tunnelId: %#v", data["tunnelId"])
|
||||
}
|
||||
if contractValueAsInt64(data["forwardCount"]) != 1 {
|
||||
t.Fatalf("expected forwardCount=1, got %#v", data["forwardCount"])
|
||||
}
|
||||
|
||||
samples, ok := data["sampleForwards"].([]interface{})
|
||||
if !ok || len(samples) != 1 {
|
||||
t.Fatalf("expected one sample forward, got %#v", data["sampleForwards"])
|
||||
}
|
||||
first, ok := samples[0].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected sample object, got %T", samples[0])
|
||||
}
|
||||
if first["name"] != "preview-forward" {
|
||||
t.Fatalf("unexpected sample name: %#v", first["name"])
|
||||
}
|
||||
if contractValueAsInt64(first["inPort"]) != 21001 {
|
||||
t.Fatalf("unexpected sample inPort: %#v", first["inPort"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelDeleteWithForwardsDeleteActionRemovesTunnelAndRulesContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
sourceTunnelID, sourceNodeID := seedTunnelDeleteTunnelWithNode(t, repo, now, "delete-source-tunnel", "delete-source-node", "22000-22010")
|
||||
forwardID := seedTunnelDeleteForward(t, repo, now, sourceTunnelID, sourceNodeID, "delete-forward", 22001)
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/tunnel/delete-with-forwards", map[string]interface{}{
|
||||
"id": sourceTunnelID,
|
||||
"action": "delete_forwards",
|
||||
})
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected success, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
if count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM tunnel WHERE id = ?`, sourceTunnelID); count != 0 {
|
||||
t.Fatalf("expected tunnel deleted, got count=%d", count)
|
||||
}
|
||||
if count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM forward WHERE id = ?`, forwardID); count != 0 {
|
||||
t.Fatalf("expected forward deleted, got count=%d", count)
|
||||
}
|
||||
if count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM forward_port WHERE forward_id = ?`, forwardID); count != 0 {
|
||||
t.Fatalf("expected forward ports deleted, got count=%d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelDeleteWithForwardsReplaceReturnsFailureDetailsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
sourceTunnelID, sourceNodeID := seedTunnelDeleteTunnelWithNode(t, repo, now, "replace-source-tunnel", "replace-source-node", "23000-23010")
|
||||
forwardID := seedTunnelDeleteForward(t, repo, now, sourceTunnelID, sourceNodeID, "replace-forward", 23001)
|
||||
targetTunnelID, targetNodeID := seedTunnelDeleteTunnelWithNode(t, repo, now, "replace-target-tunnel", "replace-target-node", "23000-23010")
|
||||
seedTunnelDeleteForward(t, repo, now, targetTunnelID, targetNodeID, "occupied-forward", 23001)
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/tunnel/delete-with-forwards", map[string]interface{}{
|
||||
"id": sourceTunnelID,
|
||||
"action": "replace",
|
||||
"targetTunnelId": targetTunnelID,
|
||||
})
|
||||
if out.Code != -2 {
|
||||
t.Fatalf("expected failure code -2, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
result := mustTunnelDeleteFailureResult(t, out)
|
||||
if contractValueAsInt64(result["failCount"]) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %#v", result["failCount"])
|
||||
}
|
||||
assertBatchFailureNameAndReason(t, result, "replace-forward", "节点 replace-target-node 端口 23001 已被其他转发占用")
|
||||
|
||||
if count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM tunnel WHERE id = ?`, sourceTunnelID); count != 1 {
|
||||
t.Fatalf("expected source tunnel kept, got count=%d", count)
|
||||
}
|
||||
if tunnelAfter := mustQueryInt64(t, repo, `SELECT tunnel_id FROM forward WHERE id = ?`, forwardID); tunnelAfter != sourceTunnelID {
|
||||
t.Fatalf("expected forward tunnel unchanged, got %d", tunnelAfter)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelBatchDeletePreviewIncludesTotalsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
tunnelA, nodeA := seedTunnelDeleteTunnelWithNode(t, repo, now, "batch-preview-a", "batch-preview-node-a", "24000-24010")
|
||||
tunnelB, _ := seedTunnelDeleteTunnelWithNode(t, repo, now, "batch-preview-b", "batch-preview-node-b", "24100-24110")
|
||||
seedTunnelDeleteForward(t, repo, now, tunnelA, nodeA, "batch-preview-forward", 24001)
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/tunnel/batch-delete-preview", map[string]interface{}{
|
||||
"ids": []int64{tunnelA, tunnelB},
|
||||
})
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected success, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
data, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected preview object, got %T", out.Data)
|
||||
}
|
||||
if contractValueAsInt64(data["tunnelCount"]) != 2 {
|
||||
t.Fatalf("expected tunnelCount=2, got %#v", data["tunnelCount"])
|
||||
}
|
||||
if contractValueAsInt64(data["totalForwardCount"]) != 1 {
|
||||
t.Fatalf("expected totalForwardCount=1, got %#v", data["totalForwardCount"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelBatchDeleteWithForwardsReturnsTunnelLevelFailuresContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
adminToken := mustAdminToken(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
sourceTunnelA, _ := seedTunnelDeleteTunnelWithNode(t, repo, now, "batch-replace-source-a", "batch-replace-source-node-a", "25000-25010")
|
||||
sourceTunnelB, sourceNodeB := seedTunnelDeleteTunnelWithNode(t, repo, now, "batch-replace-source-b", "batch-replace-source-node-b", "25100-25110")
|
||||
targetTunnelID, targetNodeID := seedTunnelDeleteTunnelWithNode(t, repo, now, "batch-replace-target", "batch-replace-target-node", "25000-25010")
|
||||
|
||||
seedTunnelDeleteForward(t, repo, now, sourceTunnelB, sourceNodeB, "batch-replace-forward-b", 25002)
|
||||
seedTunnelDeleteForward(t, repo, now, targetTunnelID, targetNodeID, "batch-replace-occupied", 25002)
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/tunnel/batch-delete-with-forwards", map[string]interface{}{
|
||||
"ids": []int64{sourceTunnelA, sourceTunnelB},
|
||||
"action": "replace",
|
||||
"targetTunnelId": targetTunnelID,
|
||||
})
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected success envelope, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
result := mustTunnelDeleteFailureResult(t, out)
|
||||
if contractValueAsInt64(result["successCount"]) != 1 {
|
||||
t.Fatalf("expected successCount=1, got %#v", result["successCount"])
|
||||
}
|
||||
if contractValueAsInt64(result["failCount"]) != 1 {
|
||||
t.Fatalf("expected failCount=1, got %#v", result["failCount"])
|
||||
}
|
||||
assertBatchFailureNameAndReason(t, result, "batch-replace-source-b", "batch-replace-forward-b: 节点 batch-replace-target-node 端口 25002 已被其他转发占用")
|
||||
|
||||
if count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM tunnel WHERE id = ?`, sourceTunnelA); count != 0 {
|
||||
t.Fatalf("expected source tunnel A deleted, got count=%d", count)
|
||||
}
|
||||
if count := mustQueryInt(t, repo, `SELECT COUNT(1) FROM tunnel WHERE id = ?`, sourceTunnelB); count != 1 {
|
||||
t.Fatalf("expected source tunnel B kept, got count=%d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func seedTunnelDeleteTunnelWithNode(t *testing.T, repo *storeRepo.Repository, now int64, tunnelName, nodeName, portRange string) (int64, int64) {
|
||||
t.Helper()
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, status, created_time, updated_time, in_ip, inx, ip_preference)
|
||||
VALUES(?, 1.0, 1, 'tls', 1, 1, ?, ?, NULL, 0, '')
|
||||
`, tunnelName, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel %s: %v", tunnelName, err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, tunnelName)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, '10.0.0.1', '10.0.0.1', '', ?, '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
|
||||
`, nodeName, nodeName+"-secret", portRange, now, now).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", nodeName, err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, nodeName)
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 0, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel for %s: %v", tunnelName, err)
|
||||
}
|
||||
|
||||
return tunnelID, nodeID
|
||||
}
|
||||
|
||||
func seedTunnelDeleteForward(t *testing.T, repo *storeRepo.Repository, now int64, tunnelID, nodeID int64, forwardName string, port int) int64 {
|
||||
t.Helper()
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(2, 'contract-user', ?, ?, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, forwardName, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward %s: %v", forwardName, err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, repo, forwardName)
|
||||
|
||||
if err := repo.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port).Error; err != nil {
|
||||
t.Fatalf("insert forward_port for %s: %v", forwardName, err)
|
||||
}
|
||||
|
||||
return forwardID
|
||||
}
|
||||
|
||||
func mustTunnelDeleteFailureResult(t *testing.T, out response.R) map[string]interface{} {
|
||||
t.Helper()
|
||||
result, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected result object, got %T", out.Data)
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestFlowUploadInsertsTunnelMetrics(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
node := &model.Node{
|
||||
Name: "node-1",
|
||||
Secret: "node-secret",
|
||||
ServerIP: "127.0.0.1",
|
||||
Port: "10000-10010",
|
||||
TCPListenAddr: "[::]",
|
||||
UDPListenAddr: "[::]",
|
||||
CreatedTime: now,
|
||||
Status: 1,
|
||||
}
|
||||
if err := repo.DB().Create(node).Error; err != nil {
|
||||
t.Fatalf("seed node: %v", err)
|
||||
}
|
||||
|
||||
tunnel := &model.Tunnel{
|
||||
Name: "tunnel-1",
|
||||
TrafficRatio: 1.0,
|
||||
Type: 1,
|
||||
Protocol: "tls",
|
||||
Flow: 1,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
}
|
||||
if err := repo.DB().Create(tunnel).Error; err != nil {
|
||||
t.Fatalf("seed tunnel: %v", err)
|
||||
}
|
||||
|
||||
forward := &model.Forward{
|
||||
UserID: 123,
|
||||
UserName: "user-123",
|
||||
Name: "forward-1",
|
||||
TunnelID: tunnel.ID,
|
||||
RemoteAddr: "1.1.1.1:80",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
}
|
||||
if err := repo.DB().Create(forward).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
|
||||
serviceName := jsonNumber(forward.ID) + "_123_0"
|
||||
body, _ := json.Marshal([]map[string]interface{}{{
|
||||
"n": serviceName,
|
||||
"u": 200,
|
||||
"d": 100,
|
||||
}})
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", res.Code)
|
||||
}
|
||||
|
||||
metrics, err := repo.GetTunnelMetrics(tunnel.ID, 0, now+60_000)
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 {
|
||||
t.Fatalf("expected 1 tunnel metric row, got %d", len(metrics))
|
||||
}
|
||||
|
||||
if metrics[0].TunnelID != tunnel.ID {
|
||||
t.Fatalf("expected tunnelId %d, got %d", tunnel.ID, metrics[0].TunnelID)
|
||||
}
|
||||
if metrics[0].NodeID != node.ID {
|
||||
t.Fatalf("expected nodeId %d, got %d", node.ID, metrics[0].NodeID)
|
||||
}
|
||||
if metrics[0].BytesIn != 100 {
|
||||
t.Fatalf("expected bytesIn 100, got %d", metrics[0].BytesIn)
|
||||
}
|
||||
if metrics[0].BytesOut != 200 {
|
||||
t.Fatalf("expected bytesOut 200, got %d", metrics[0].BytesOut)
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package contract
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
@@ -13,7 +13,7 @@ import (
|
||||
|
||||
func TestUserTunnelVisibleListContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupDiagnosisContractRouter(t, secret)
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
@@ -126,8 +126,11 @@ func collectTunnelIDs(t *testing.T, data interface{}) map[int64]bool {
|
||||
if !ok {
|
||||
t.Fatalf("expected object item, got %T", item)
|
||||
}
|
||||
id := int64(obj["id"].(float64))
|
||||
ids[id] = true
|
||||
idFloat, ok := obj["id"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected id to be float64, got %T", obj["id"])
|
||||
}
|
||||
ids[int64(idFloat)] = true
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestForwardCreateBlockedWhenUserQuotaExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'quota_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'quota_tunnel', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(2, "quota_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewBufferString(`{"tunnelId":1,"name":"quota-forward","remoteAddr":"1.1.1.1:53"}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when user quota exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "配额") {
|
||||
t.Fatalf("expected quota error, got %q", out.Msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResumeBlockedWhenUserQuotaExceeded(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'quota_resume_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(1, 'quota_resume_tunnel', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := repo.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(1, 2, 'quota_resume_user', 'quota_resume_forward', 1, '1.1.1.1:53', 'fifo', 0, 0, ?, ?, 0, 0)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '1', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(2, "quota_resume_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/resume", bytes.NewBufferString(`{"id":1}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected non-zero code when user quota exceeded")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "配额") {
|
||||
t.Fatalf("expected quota error, got %q", out.Msg)
|
||||
}
|
||||
status := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = 1`)
|
||||
if status != 0 {
|
||||
t.Fatalf("expected forward to remain paused, got %d", status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserQuotaResetClearsDisableFlag(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'quota_reset_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
|
||||
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '', ?, ?)
|
||||
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
token, err := auth.GenerateToken(1, "admin", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/quota/reset", bytes.NewBufferString(`{"userId":2,"scope":"all"}`))
|
||||
req.Header.Set("Authorization", token)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected reset success, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
quotaDisabled := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`)
|
||||
if quotaDisabled != 0 {
|
||||
t.Fatalf("expected quota disable flag cleared, got %d", quotaDisabled)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestUserTunnelListReturnsStoredStatusContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(201, 'user_tunnel_status_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(301, 'user-tunnel-status-enabled', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel enabled: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(302, 'user-tunnel-status-disabled', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel disabled: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(401, 201, 301, NULL, 10, 500, 0, 0, 1, 2727251700000, 1)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert enabled user_tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(402, 201, 302, NULL, 10, 500, 0, 0, 1, 2727251700000, 0)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert disabled user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
body := bytes.NewBufferString(`{"userId":201}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/list", body)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d (%s)", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
items, ok := out.Data.([]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected array data, got %T", out.Data)
|
||||
}
|
||||
if len(items) != 2 {
|
||||
t.Fatalf("expected 2 items, got %d", len(items))
|
||||
}
|
||||
|
||||
statusByTunnelID := make(map[int64]int, len(items))
|
||||
for _, item := range items {
|
||||
obj, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected object item, got %T", item)
|
||||
}
|
||||
tunnelID, ok := obj["tunnelId"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected tunnelId to be float64, got %T", obj["tunnelId"])
|
||||
}
|
||||
status, ok := obj["status"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected status to be float64, got %T", obj["status"])
|
||||
}
|
||||
statusByTunnelID[int64(tunnelID)] = int(status)
|
||||
}
|
||||
|
||||
if statusByTunnelID[301] != 1 {
|
||||
t.Fatalf("expected enabled tunnel status 1, got %d", statusByTunnelID[301])
|
||||
}
|
||||
if statusByTunnelID[302] != 0 {
|
||||
t.Fatalf("expected disabled tunnel status 0, got %d", statusByTunnelID[302])
|
||||
}
|
||||
}
|
||||
+10
-5
@@ -1,6 +1,9 @@
|
||||
# GO-GOST SERVICE KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Forwarding agent built on GOST v3 with a local fork of `github.com/go-gost/x` under `x/`.
|
||||
@@ -19,16 +22,18 @@ go-gost/
|
||||
## WHERE TO LOOK
|
||||
| Task | Location | Notes |
|
||||
|------|----------|-------|
|
||||
| Panel integration config | `go-gost/config.go` | Expects `config.json` in cwd by default |
|
||||
| Service lifecycle/reload | `go-gost/program.go` | Parses config; handles SIGHUP reload |
|
||||
| WebSocket reporting | `go-gost/main.go` | Starts reporter + sets HTTP report URL |
|
||||
| Protocol behaviors | `go-gost/x/` | Handlers/listeners/dialers live here |
|
||||
| **Panel integration config** | `go-gost/config.go` | Expects `config.json` in cwd by default |
|
||||
| **Service lifecycle/reload** | `go-gost/program.go` | Parses config; handles SIGHUP reload |
|
||||
| **WebSocket reporting** | `go-gost/main.go` | Starts reporter + sets HTTP report URL |
|
||||
| **Protocol behaviors** | `go-gost/x/` | Handlers/listeners/dialers live here |
|
||||
| **Build** | `go-gost/Makefile` | Cross-compile targets for amd64/arm64 |
|
||||
|
||||
## CONVENTIONS
|
||||
- Two configs exist: panel integration uses `config.json`; forwarding services use GOST config (defaults to `gost.{json,yaml}` via viper search paths).
|
||||
- `go-gost/x/` is the primary extension surface; avoid editing vendored deps.
|
||||
- Agent communicates with panel via WebSocket (real-time commands) + HTTP (batch traffic reports).
|
||||
- All panel communication uses AES encryption with node `secret` as PSK.
|
||||
- CI builds with `CGO_ENABLED=0` for static binaries, then compresses with UPX.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- **DO NOT EDIT** generated protobuf in `x/internal/util/grpc/proto/`.
|
||||
|
||||
+5
-3
@@ -109,17 +109,19 @@ func main() {
|
||||
// 加载配置文件
|
||||
config, err := LoadConfig("config.json")
|
||||
if err != nil {
|
||||
fmt.Println("❌ 配置加载失败: %v\n", err)
|
||||
fmt.Printf("❌ 配置加载失败: %v\n", err)
|
||||
fmt.Println("请确保当前目录存在 config.json 文件")
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
fmt.Println("✅ 配置加载成功 - addr: %s", config.Addr)
|
||||
fmt.Printf("✅ 配置加载成功 - addr: %s\n", config.Addr)
|
||||
|
||||
log := xlogger.NewLogger()
|
||||
logger.SetDefault(log)
|
||||
|
||||
wsReporter := socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret, config.Http, config.Tls, config.Socks, version)
|
||||
distro := socket.DetectDistro()
|
||||
fullVersion := fmt.Sprintf("%s (%s/%s)", version, distro, runtime.GOARCH)
|
||||
wsReporter := socket.StartWebSocketReporterWithConfig(config.Addr, config.Secret, config.Http, config.Tls, config.Socks, fullVersion)
|
||||
defer wsReporter.Stop()
|
||||
service.SetHTTPReportURL(config.Addr, config.Secret)
|
||||
|
||||
|
||||
+16
-10
@@ -1,32 +1,38 @@
|
||||
# GO-GOST/X KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Local fork of `github.com/go-gost/x` used by `go-gost/` via `replace github.com/go-gost/x => ./x`. Most protocol/runtime behavior changes happen here. 30+ top-level packages - framework-style layout.
|
||||
|
||||
## STRUCTURE
|
||||
```
|
||||
go-gost/x/
|
||||
├── api/ # Gin management API + embedded swagger docs
|
||||
├── api/ # Gin management API + embedded swagger docs (22 files)
|
||||
├── config/ # Config model + parsing/load/reload
|
||||
├── connector/ # Outbound connect implementations
|
||||
├── dialer/ # Outbound dialers (tcp/tls/ws/quic/...)
|
||||
├── dialer/ # Outbound dialers (tcp/tls/ws/quic/...)
|
||||
├── handler/ # Protocol handlers (socks/http/tunnel/relay/...)
|
||||
├── listener/ # Inbound listeners (tcp/udp/tun/tap/redirect/...)
|
||||
├── limiter/ # Traffic/rate/conn limiters
|
||||
├── registry/ # Registries for services/handlers/listeners/etc
|
||||
├── registry/ # Registries for services/handlers/listeners/etc (20 files)
|
||||
├── service/ # Service wrappers + reporting hooks
|
||||
├── socket/ # WebSocket reporter / panel integration
|
||||
├── socket/ # WebSocket reporter / panel integration (6 files)
|
||||
└── internal/ # Shared internals (grpc proto, net utils, sniffing, tls, ...)
|
||||
```
|
||||
|
||||
## WHERE TO LOOK
|
||||
| Task | Location | Notes |
|
||||
|------|----------|-------|
|
||||
| Management API routes/auth | `go-gost/x/api/api.go` | `/docs`, `/config/*`; BasicAuth + interceptor |
|
||||
| Service config parsing | `go-gost/x/config/parsing/` | Converts config to running services |
|
||||
| Add a handler | `go-gost/x/handler/` | Per-protocol subdirs |
|
||||
| Add a listener/dialer | `go-gost/x/listener/`, `go-gost/x/dialer/` | Transport variants |
|
||||
| Panel reporting | `go-gost/x/socket/` | WebSocket + HTTP report URL hooks |
|
||||
| **Management API routes/auth** | `go-gost/x/api/api.go` | `/docs`, `/config/*`; BasicAuth + interceptor |
|
||||
| **Service config parsing** | `go-gost/x/config/parsing/` | Converts config to running services |
|
||||
| **Add a handler** | `go-gost/x/handler/` | Per-protocol subdirs |
|
||||
| **Add a listener/dialer** | `go-gost/x/listener/`, `go-gost/x/dialer/` | Transport variants |
|
||||
| **Panel reporting** | `go-gost/x/socket/` | WebSocket + HTTP report URL hooks |
|
||||
| **Register new component** | `go-gost/x/registry/` | `Register{Type}(name, creator)` |
|
||||
|
||||
## CONVENTIONS
|
||||
- `go-gost/x/` is a standalone Go module (`go-gost/x/go.mod`); run go tooling from this dir when debugging module resolution.
|
||||
@@ -41,4 +47,4 @@ go-gost/x/
|
||||
```bash
|
||||
cd go-gost/x
|
||||
go test ./...
|
||||
```
|
||||
```
|
||||
@@ -1,5 +1,10 @@
|
||||
# GO-GOST/X API KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Gin-based management API for reading/writing config and controlling services at runtime.
|
||||
|
||||
|
||||
@@ -624,6 +624,11 @@ func resumeService(ctx *gin.Context) {
|
||||
existingSvc.Close()
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
|
||||
// 强制断开端口的所有连接
|
||||
if serviceConfig.Addr != "" {
|
||||
_ = kill.ForceClosePortConnections(serviceConfig.Addr)
|
||||
}
|
||||
|
||||
// 等待端口释放
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
@@ -1039,8 +1044,13 @@ func resumeServices(ctx *gin.Context) {
|
||||
str.service.Close()
|
||||
registry.ServiceRegistry().Unregister(str.name)
|
||||
|
||||
// 强制断开端口的所有连接
|
||||
if str.serviceConfig.Addr != "" {
|
||||
_ = kill.ForceClosePortConnections(str.serviceConfig.Addr)
|
||||
}
|
||||
|
||||
// 等待端口释放
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
// 重新解析并启动服务
|
||||
svc, err := parser.ParseService(str.serviceConfig)
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
# GO-GOST/X CONFIG KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Config model + parsing/loading pipeline for the `go-gost/x` runtime. This is the bridge between `gost.json`/`gost.yaml` and in-memory registries/services.
|
||||
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
# GOST CONNECTOR KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Feb 13 2026
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Connection initiators (clients) for various protocols in GOST forwarding.
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
# GO-GOST/X DIALERS KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Outbound dialers (client-side connection establishment) used by connectors/handlers.
|
||||
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
# GO-GOST/X HANDLERS KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Protocol handlers (server-side request handling) used by services defined in the GOST config.
|
||||
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
# GO-GOST/X LISTENERS KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Inbound listeners (transport-level accept loops) used by services defined in the GOST config.
|
||||
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
# GO-GOST REGISTRY KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Wed Feb 04 2026
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
Central registration point for all pluggable GOST components (handlers, listeners, dialers, etc.).
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
@@ -63,11 +62,9 @@ func SetProtocolBlock(httpOn int, tlsOn int, socksOn int) {
|
||||
type Option func(opts *options)
|
||||
|
||||
func init() {
|
||||
_, err := LoadConfig("config.json")
|
||||
fmt.Println("config.json loaded")
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
// NOTE: This package can be imported by tests/tools that don't have a local
|
||||
// config.json. Missing config should not crash the process.
|
||||
_, _ = LoadConfig("config.json")
|
||||
needWrap = isTls+isSocks+isHttp > 0
|
||||
}
|
||||
|
||||
|
||||
@@ -6,7 +6,9 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/go-gost/core/observer/stats"
|
||||
@@ -18,6 +20,15 @@ import (
|
||||
var httpReportURL string
|
||||
var configReportURL string
|
||||
var httpAESCrypto *crypto.AESCrypto // 新增:HTTP上报加密器
|
||||
var reportURLPreferenceMutex sync.RWMutex
|
||||
var preferredUploadURL string
|
||||
var preferredConfigURL string
|
||||
var reportDo = func(ctx context.Context, req *http.Request, timeout time.Duration) (*http.Response, error) {
|
||||
client := &http.Client{
|
||||
Timeout: timeout,
|
||||
}
|
||||
return client.Do(req.WithContext(ctx))
|
||||
}
|
||||
|
||||
// TrafficReportItem 流量报告项(压缩格式)
|
||||
type TrafficReportItem struct {
|
||||
@@ -27,8 +38,17 @@ type TrafficReportItem struct {
|
||||
}
|
||||
|
||||
func SetHTTPReportURL(addr string, secret string) {
|
||||
httpReportURL = "http://" + addr + "/flow/upload?secret=" + secret
|
||||
configReportURL = "http://" + addr + "/flow/config?secret=" + secret
|
||||
uploadURLs, configURLs := buildReportURLCandidates(addr, secret)
|
||||
if len(uploadURLs) > 0 {
|
||||
httpReportURL = strings.Join(uploadURLs, ",")
|
||||
}
|
||||
if len(configURLs) > 0 {
|
||||
configReportURL = strings.Join(configURLs, ",")
|
||||
}
|
||||
reportURLPreferenceMutex.Lock()
|
||||
preferredUploadURL = ""
|
||||
preferredConfigURL = ""
|
||||
reportURLPreferenceMutex.Unlock()
|
||||
|
||||
// 创建 AES 加密器
|
||||
var err error
|
||||
@@ -41,8 +61,173 @@ func SetHTTPReportURL(addr string, secret string) {
|
||||
}
|
||||
}
|
||||
|
||||
func buildReportURLCandidates(addr string, secret string) (upload []string, config []string) {
|
||||
normalizedAddr, explicitScheme := normalizeReportAddress(addr)
|
||||
if normalizedAddr == "" {
|
||||
normalizedAddr = strings.TrimSpace(addr)
|
||||
}
|
||||
|
||||
schemes := []string{"https", "http"}
|
||||
if mappedScheme := mapToHTTPScheme(explicitScheme); mappedScheme == "http" {
|
||||
schemes = []string{"http", "https"}
|
||||
}
|
||||
|
||||
upload = []string{
|
||||
schemes[0] + "://" + normalizedAddr + "/flow/upload?secret=" + secret,
|
||||
schemes[1] + "://" + normalizedAddr + "/flow/upload?secret=" + secret,
|
||||
}
|
||||
config = []string{
|
||||
schemes[0] + "://" + normalizedAddr + "/flow/config?secret=" + secret,
|
||||
schemes[1] + "://" + normalizedAddr + "/flow/config?secret=" + secret,
|
||||
}
|
||||
return upload, config
|
||||
}
|
||||
|
||||
func normalizeReportAddress(addr string) (string, string) {
|
||||
raw := strings.TrimSpace(addr)
|
||||
if raw == "" {
|
||||
return "", ""
|
||||
}
|
||||
|
||||
scheme := ""
|
||||
if idx := strings.Index(raw, "://"); idx > 0 {
|
||||
scheme = strings.ToLower(strings.TrimSpace(raw[:idx]))
|
||||
if parsed, err := url.Parse(raw); err == nil {
|
||||
if host := strings.TrimSpace(parsed.Host); host != "" {
|
||||
return host, scheme
|
||||
}
|
||||
}
|
||||
raw = raw[idx+3:]
|
||||
}
|
||||
|
||||
if idx := strings.IndexAny(raw, "/?#"); idx >= 0 {
|
||||
raw = raw[:idx]
|
||||
}
|
||||
return strings.TrimSpace(raw), scheme
|
||||
}
|
||||
|
||||
func mapToHTTPScheme(scheme string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(scheme)) {
|
||||
case "https", "wss":
|
||||
return "https"
|
||||
case "http", "ws":
|
||||
return "http"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func loadPreferredURL(preferred *string) string {
|
||||
if preferred == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
reportURLPreferenceMutex.RLock()
|
||||
defer reportURLPreferenceMutex.RUnlock()
|
||||
return *preferred
|
||||
}
|
||||
|
||||
func storePreferredURL(preferred *string, value string) {
|
||||
if preferred == nil {
|
||||
return
|
||||
}
|
||||
|
||||
reportURLPreferenceMutex.Lock()
|
||||
defer reportURLPreferenceMutex.Unlock()
|
||||
*preferred = value
|
||||
}
|
||||
|
||||
func prioritizeURLs(urls []string, preferred string) []string {
|
||||
ordered := append([]string(nil), urls...)
|
||||
if preferred == "" || len(ordered) < 2 {
|
||||
return ordered
|
||||
}
|
||||
|
||||
for i, targetURL := range ordered {
|
||||
if targetURL == preferred {
|
||||
if i > 0 {
|
||||
ordered[0], ordered[i] = ordered[i], ordered[0]
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return ordered
|
||||
}
|
||||
|
||||
func postJSONWithFallback(ctx context.Context, urls []string, requestBody []byte, userAgent string, timeout time.Duration, preferred *string) (bool, error) {
|
||||
if len(urls) == 0 {
|
||||
return false, fmt.Errorf("上报URL未设置")
|
||||
}
|
||||
|
||||
orderedURLs := prioritizeURLs(urls, loadPreferredURL(preferred))
|
||||
|
||||
var errs []string
|
||||
for i, targetURL := range orderedURLs {
|
||||
req, err := http.NewRequest("POST", targetURL, bytes.NewBuffer(requestBody))
|
||||
if err != nil {
|
||||
errs = append(errs, fmt.Sprintf("%s => 创建请求失败: %v", targetURL, err))
|
||||
if i < len(orderedURLs)-1 {
|
||||
fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 创建请求失败: %v\n", targetURL, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("User-Agent", userAgent)
|
||||
|
||||
resp, err := reportDo(ctx, req, timeout)
|
||||
if err != nil {
|
||||
errs = append(errs, fmt.Sprintf("%s => 请求失败: %v", targetURL, err))
|
||||
if i < len(orderedURLs)-1 {
|
||||
fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 请求失败: %v\n", targetURL, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
var responseBytes bytes.Buffer
|
||||
_, readErr := responseBytes.ReadFrom(resp.Body)
|
||||
resp.Body.Close()
|
||||
if readErr != nil {
|
||||
errs = append(errs, fmt.Sprintf("%s => 读取响应失败: %v", targetURL, readErr))
|
||||
if i < len(orderedURLs)-1 {
|
||||
fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 读取响应失败: %v\n", targetURL, readErr)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
errs = append(errs, fmt.Sprintf("%s => HTTP响应错误: %d %s", targetURL, resp.StatusCode, resp.Status))
|
||||
if i < len(orderedURLs)-1 {
|
||||
fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => HTTP响应错误: %d %s\n", targetURL, resp.StatusCode, resp.Status)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
responseText := strings.TrimSpace(responseBytes.String())
|
||||
if responseText == "ok" {
|
||||
if i > 0 {
|
||||
fmt.Printf("↪️ HTTP上报已自动回退到: %s\n", targetURL)
|
||||
}
|
||||
storePreferredURL(preferred, targetURL)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
errs = append(errs, fmt.Sprintf("%s => 服务器响应: %s (期望: ok)", targetURL, responseText))
|
||||
if i < len(orderedURLs)-1 {
|
||||
fmt.Printf("⚠️ HTTP上报尝试失败,准备回退: %s => 服务器响应: %s (期望: ok)\n", targetURL, responseText)
|
||||
}
|
||||
}
|
||||
|
||||
return false, fmt.Errorf("发送HTTP请求失败: %s", strings.Join(errs, " | "))
|
||||
}
|
||||
|
||||
// sendBatchTrafficReport 批量发送多个服务的流量报告到HTTP接口
|
||||
func sendBatchTrafficReport(ctx context.Context, reportItems []TrafficReportItem) (bool, error) {
|
||||
if httpReportURL == "" {
|
||||
return false, fmt.Errorf("流量上报URL未设置")
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reportItems)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("序列化报告数据失败: %v", err)
|
||||
@@ -73,46 +258,16 @@ func sendBatchTrafficReport(ctx context.Context, reportItems []TrafficReportItem
|
||||
requestBody = jsonData
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", httpReportURL, bytes.NewBuffer(requestBody))
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("创建HTTP请求失败: %v", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("User-Agent", "GOST-Traffic-Reporter/1.0")
|
||||
|
||||
client := &http.Client{
|
||||
Timeout: 5 * time.Second,
|
||||
}
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("发送HTTP请求失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return false, fmt.Errorf("HTTP响应错误: %d %s", resp.StatusCode, resp.Status)
|
||||
}
|
||||
|
||||
// 读取响应内容
|
||||
var responseBytes bytes.Buffer
|
||||
_, err = responseBytes.ReadFrom(resp.Body)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("读取响应内容失败: %v", err)
|
||||
}
|
||||
|
||||
responseText := strings.TrimSpace(responseBytes.String())
|
||||
|
||||
// 检查响应是否为"ok"
|
||||
if responseText == "ok" {
|
||||
return true, nil
|
||||
} else {
|
||||
return false, fmt.Errorf("服务器响应: %s (期望: ok)", responseText)
|
||||
}
|
||||
return postJSONWithFallback(
|
||||
ctx,
|
||||
strings.Split(httpReportURL, ","),
|
||||
requestBody,
|
||||
"GOST-Traffic-Reporter/1.0",
|
||||
5*time.Second,
|
||||
&preferredUploadURL,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
// sendConfigReport 发送配置报告到HTTP接口
|
||||
func sendConfigReport(ctx context.Context) (bool, error) {
|
||||
if configReportURL == "" {
|
||||
@@ -150,43 +305,14 @@ func sendConfigReport(ctx context.Context) (bool, error) {
|
||||
requestBody = configData
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", configReportURL, bytes.NewBuffer(requestBody))
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("创建HTTP请求失败: %v", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("User-Agent", "Config-Reporter/1.0")
|
||||
|
||||
client := &http.Client{
|
||||
Timeout: 10 * time.Second, // 配置上报可以稍长一些
|
||||
}
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("发送HTTP请求失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return false, fmt.Errorf("HTTP响应错误: %d %s", resp.StatusCode, resp.Status)
|
||||
}
|
||||
|
||||
// 读取响应内容
|
||||
var responseBytes bytes.Buffer
|
||||
_, err = responseBytes.ReadFrom(resp.Body)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("读取响应内容失败: %v", err)
|
||||
}
|
||||
|
||||
responseText := strings.TrimSpace(responseBytes.String())
|
||||
|
||||
// 检查响应是否为"ok"
|
||||
if responseText == "ok" {
|
||||
return true, nil
|
||||
} else {
|
||||
return false, fmt.Errorf("服务器响应: %s (期望: ok)", responseText)
|
||||
}
|
||||
return postJSONWithFallback(
|
||||
ctx,
|
||||
strings.Split(configReportURL, ","),
|
||||
requestBody,
|
||||
"Config-Reporter/1.0",
|
||||
10*time.Second,
|
||||
&preferredConfigURL,
|
||||
)
|
||||
}
|
||||
|
||||
// StartConfigReporter 启动配置定时上报器(每10分钟上报一次)
|
||||
|
||||
@@ -0,0 +1,150 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestBuildReportURLCandidatesSecureFirst(t *testing.T) {
|
||||
upload, config := buildReportURLCandidates("panel.example.com:443", "abc")
|
||||
|
||||
if len(upload) != 2 {
|
||||
t.Fatalf("expected 2 upload candidates, got %d", len(upload))
|
||||
}
|
||||
if len(config) != 2 {
|
||||
t.Fatalf("expected 2 config candidates, got %d", len(config))
|
||||
}
|
||||
|
||||
if upload[0] != "https://panel.example.com:443/flow/upload?secret=abc" {
|
||||
t.Fatalf("unexpected upload[0]: %s", upload[0])
|
||||
}
|
||||
if upload[1] != "http://panel.example.com:443/flow/upload?secret=abc" {
|
||||
t.Fatalf("unexpected upload[1]: %s", upload[1])
|
||||
}
|
||||
if config[0] != "https://panel.example.com:443/flow/config?secret=abc" {
|
||||
t.Fatalf("unexpected config[0]: %s", config[0])
|
||||
}
|
||||
if config[1] != "http://panel.example.com:443/flow/config?secret=abc" {
|
||||
t.Fatalf("unexpected config[1]: %s", config[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildReportURLCandidatesNormalizeSchemeAddr(t *testing.T) {
|
||||
upload, config := buildReportURLCandidates("https://panel.example.com:8443/path", "abc")
|
||||
|
||||
if upload[0] != "https://panel.example.com:8443/flow/upload?secret=abc" {
|
||||
t.Fatalf("unexpected upload[0]: %s", upload[0])
|
||||
}
|
||||
if upload[1] != "http://panel.example.com:8443/flow/upload?secret=abc" {
|
||||
t.Fatalf("unexpected upload[1]: %s", upload[1])
|
||||
}
|
||||
if config[0] != "https://panel.example.com:8443/flow/config?secret=abc" {
|
||||
t.Fatalf("unexpected config[0]: %s", config[0])
|
||||
}
|
||||
if config[1] != "http://panel.example.com:8443/flow/config?secret=abc" {
|
||||
t.Fatalf("unexpected config[1]: %s", config[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostJSONWithFallbackUsesHTTPAfterHTTPSFailure(t *testing.T) {
|
||||
orig := reportDo
|
||||
defer func() { reportDo = orig }()
|
||||
|
||||
var calls []string
|
||||
reportDo = func(_ context.Context, req *http.Request, _ time.Duration) (*http.Response, error) {
|
||||
calls = append(calls, req.URL.String())
|
||||
if strings.HasPrefix(req.URL.String(), "https://") {
|
||||
return nil, errors.New("tls handshake failed")
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader("ok")),
|
||||
}, nil
|
||||
}
|
||||
|
||||
ok, err := postJSONWithFallback(
|
||||
context.Background(),
|
||||
[]string{
|
||||
"https://panel.example.com:443/flow/upload?secret=abc",
|
||||
"http://panel.example.com:443/flow/upload?secret=abc",
|
||||
},
|
||||
[]byte(`[]`),
|
||||
"GOST-Traffic-Reporter/1.0",
|
||||
5*time.Second,
|
||||
nil,
|
||||
)
|
||||
if !ok || err != nil {
|
||||
t.Fatalf("expected fallback success, ok=%v err=%v", ok, err)
|
||||
}
|
||||
if len(calls) != 2 {
|
||||
t.Fatalf("expected 2 calls, got %d", len(calls))
|
||||
}
|
||||
if !strings.HasPrefix(calls[0], "https://") || !strings.HasPrefix(calls[1], "http://") {
|
||||
t.Fatalf("unexpected call order: %#v", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostJSONWithFallbackRemembersDetectedURL(t *testing.T) {
|
||||
orig := reportDo
|
||||
defer func() { reportDo = orig }()
|
||||
|
||||
targets := []string{
|
||||
"https://panel.example.com:443/flow/upload?secret=abc",
|
||||
"http://panel.example.com:443/flow/upload?secret=abc",
|
||||
}
|
||||
|
||||
var preferred string
|
||||
var calls []string
|
||||
reportDo = func(_ context.Context, req *http.Request, _ time.Duration) (*http.Response, error) {
|
||||
calls = append(calls, req.URL.String())
|
||||
if strings.HasPrefix(req.URL.String(), "https://") {
|
||||
return nil, errors.New("tls handshake failed")
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader("ok")),
|
||||
}, nil
|
||||
}
|
||||
|
||||
ok, err := postJSONWithFallback(
|
||||
context.Background(),
|
||||
targets,
|
||||
[]byte(`[]`),
|
||||
"GOST-Traffic-Reporter/1.0",
|
||||
5*time.Second,
|
||||
&preferred,
|
||||
)
|
||||
if !ok || err != nil {
|
||||
t.Fatalf("expected first call success, ok=%v err=%v", ok, err)
|
||||
}
|
||||
if preferred != targets[1] {
|
||||
t.Fatalf("expected preferred url to be remembered as %s, got %s", targets[1], preferred)
|
||||
}
|
||||
if len(calls) != 2 {
|
||||
t.Fatalf("expected 2 calls on first attempt, got %d", len(calls))
|
||||
}
|
||||
|
||||
calls = nil
|
||||
ok, err = postJSONWithFallback(
|
||||
context.Background(),
|
||||
targets,
|
||||
[]byte(`[]`),
|
||||
"GOST-Traffic-Reporter/1.0",
|
||||
5*time.Second,
|
||||
&preferred,
|
||||
)
|
||||
if !ok || err != nil {
|
||||
t.Fatalf("expected second call success, ok=%v err=%v", ok, err)
|
||||
}
|
||||
if len(calls) != 1 {
|
||||
t.Fatalf("expected second call to use remembered url once, got %d calls", len(calls))
|
||||
}
|
||||
if !strings.HasPrefix(calls[0], "http://") {
|
||||
t.Fatalf("expected remembered http url first, got %s", calls[0])
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,9 @@
|
||||
# GOST SOCKET KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Generated:** Fri Mar 20 2026
|
||||
**Commit:** f45f960
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.9-beta6
|
||||
|
||||
## OVERVIEW
|
||||
WebSocket reporter and socket utilities for panel integration.
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/shirou/gopsutil/v3/host"
|
||||
)
|
||||
|
||||
// DetectDistro returns the Linux distribution name (e.g. "ubuntu", "centos",
|
||||
// "debian"). Falls back to "linux" when detection fails.
|
||||
func DetectDistro() string {
|
||||
info, err := host.Info()
|
||||
if err != nil || info == nil {
|
||||
return "linux"
|
||||
}
|
||||
platform := strings.ToLower(strings.TrimSpace(info.Platform))
|
||||
if platform == "" {
|
||||
return "linux"
|
||||
}
|
||||
return platform
|
||||
}
|
||||
@@ -397,8 +397,13 @@ func resumeServices(req resumeServicesRequest) error {
|
||||
str.service.Close()
|
||||
registry.ServiceRegistry().Unregister(str.name)
|
||||
|
||||
// 强制断开端口的所有连接
|
||||
if str.serviceConfig.Addr != "" {
|
||||
_ = kill.ForceClosePortConnections(str.serviceConfig.Addr)
|
||||
}
|
||||
|
||||
// 等待端口释放
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
// 重新解析并启动服务
|
||||
svc, err := parser.ParseService(str.serviceConfig)
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"math/rand"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
@@ -17,7 +18,7 @@ import (
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync" // 新增:用于管理连接状态的互斥锁
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/go-gost/x/config"
|
||||
@@ -25,34 +26,67 @@ import (
|
||||
"github.com/go-gost/x/service"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/shirou/gopsutil/v3/cpu"
|
||||
"github.com/shirou/gopsutil/v3/disk"
|
||||
"github.com/shirou/gopsutil/v3/host"
|
||||
"github.com/shirou/gopsutil/v3/load"
|
||||
"github.com/shirou/gopsutil/v3/mem"
|
||||
psnet "github.com/shirou/gopsutil/v3/net"
|
||||
"golang.org/x/net/icmp"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
// SystemInfo 系统信息结构体
|
||||
type SystemInfo struct {
|
||||
Uptime uint64 `json:"uptime"` // 开机时间 (秒)
|
||||
BytesReceived uint64 `json:"bytes_received"` // 接收字节数
|
||||
BytesTransmitted uint64 `json:"bytes_transmitted"` // 发送字节数
|
||||
CPUUsage float64 `json:"cpu_usage"` // CPU使用率(百分比)
|
||||
MemoryUsage float64 `json:"memory_usage"` // 内存使用率(百分比)
|
||||
Uptime uint64 `json:"uptime"`
|
||||
BytesReceived uint64 `json:"bytes_received"`
|
||||
BytesTransmitted uint64 `json:"bytes_transmitted"`
|
||||
CPUUsage float64 `json:"cpu_usage"`
|
||||
MemoryUsage float64 `json:"memory_usage"`
|
||||
DiskUsage float64 `json:"disk_usage"`
|
||||
Load1 float64 `json:"load1"`
|
||||
Load5 float64 `json:"load5"`
|
||||
Load15 float64 `json:"load15"`
|
||||
TCPConns int64 `json:"tcp_conns"`
|
||||
UDPConns int64 `json:"udp_conns"`
|
||||
NetInSpeed int64 `json:"net_in_speed"`
|
||||
NetOutSpeed int64 `json:"net_out_speed"`
|
||||
}
|
||||
|
||||
// NetworkStats 网络统计信息
|
||||
type NetworkStats struct {
|
||||
BytesReceived uint64 `json:"bytes_received"` // 接收字节数
|
||||
BytesTransmitted uint64 `json:"bytes_transmitted"` // 发送字节数
|
||||
BytesReceived uint64 `json:"bytes_received"`
|
||||
BytesTransmitted uint64 `json:"bytes_transmitted"`
|
||||
BytesRecvDelta uint64 `json:"bytes_recv_delta"`
|
||||
BytesSentDelta uint64 `json:"bytes_sent_delta"`
|
||||
}
|
||||
|
||||
// CPUInfo CPU信息
|
||||
type CPUInfo struct {
|
||||
Usage float64 `json:"usage"` // CPU使用率(百分比)
|
||||
Usage float64 `json:"usage"`
|
||||
}
|
||||
|
||||
// MemoryInfo 内存信息
|
||||
type MemoryInfo struct {
|
||||
Usage float64 `json:"usage"` // 内存使用率(百分比)
|
||||
Usage float64 `json:"usage"`
|
||||
}
|
||||
|
||||
// DiskInfo 磁盘信息
|
||||
type DiskInfo struct {
|
||||
Usage float64 `json:"usage"`
|
||||
}
|
||||
|
||||
// LoadInfo 负载信息
|
||||
type LoadInfo struct {
|
||||
Load1 float64 `json:"load1"`
|
||||
Load5 float64 `json:"load5"`
|
||||
Load15 float64 `json:"load15"`
|
||||
}
|
||||
|
||||
// ConnectionInfo 连接信息
|
||||
type ConnectionInfo struct {
|
||||
TCPConns int64 `json:"tcp_conns"`
|
||||
UDPConns int64 `json:"udp_conns"`
|
||||
}
|
||||
|
||||
// CommandMessage 命令消息结构体
|
||||
@@ -91,26 +125,53 @@ type TcpPingResponse struct {
|
||||
RequestId string `json:"requestId,omitempty"`
|
||||
}
|
||||
|
||||
// ServiceMonitorCheckRequest service monitor check request.
|
||||
type ServiceMonitorCheckRequest struct {
|
||||
MonitorID int64 `json:"monitorId"`
|
||||
Type string `json:"type"` // tcp|icmp
|
||||
Target string `json:"target"`
|
||||
TimeoutSec int `json:"timeoutSec"`
|
||||
}
|
||||
|
||||
// ServiceMonitorCheckResult node-executed check output.
|
||||
// CommandResponse.Success indicates command execution status.
|
||||
// Actual check success is represented by this struct.
|
||||
type ServiceMonitorCheckResult struct {
|
||||
MonitorID int64 `json:"monitorId"`
|
||||
Success bool `json:"success"`
|
||||
LatencyMs float64 `json:"latencyMs"`
|
||||
StatusCode int `json:"statusCode,omitempty"`
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
}
|
||||
|
||||
const (
|
||||
reporterReadWait = 60 * time.Second
|
||||
reporterWriteWait = 5 * time.Second
|
||||
wsPingInterval = 20 * time.Second // 独立 WebSocket ping 间隔
|
||||
initialBackoff = 2 * time.Second // 重连初始退避
|
||||
maxBackoff = 2 * time.Minute // 重连最大退避
|
||||
)
|
||||
|
||||
type WebSocketReporter struct {
|
||||
url string
|
||||
addr string // 保存服务器地址
|
||||
secret string // 保存密钥
|
||||
version string // 保存版本号
|
||||
conn *websocket.Conn
|
||||
reconnectTime time.Duration
|
||||
pingInterval time.Duration
|
||||
configInterval time.Duration
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
connected bool
|
||||
connecting bool // 新增:正在连接状态
|
||||
connMutex sync.Mutex // 新增:连接状态锁
|
||||
aesCrypto *crypto.AESCrypto // 新增:AES加密器
|
||||
url string
|
||||
addr string // 保存服务器地址
|
||||
secret string // 保存密钥
|
||||
version string // 保存版本号
|
||||
preferredWSScheme string
|
||||
conn *websocket.Conn
|
||||
curBackoff time.Duration // 当前重连退避间隔
|
||||
pingInterval time.Duration
|
||||
configInterval time.Duration
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
connected bool
|
||||
connecting bool // 正在连接状态
|
||||
connMutex sync.Mutex // 连接状态锁
|
||||
aesCrypto *crypto.AESCrypto // AES加密器
|
||||
}
|
||||
|
||||
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||
return dialer.Dial(rawURL, nil)
|
||||
}
|
||||
|
||||
// NewWebSocketReporter 创建一个新的WebSocket报告器
|
||||
@@ -128,8 +189,8 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
|
||||
|
||||
return &WebSocketReporter{
|
||||
url: serverURL,
|
||||
reconnectTime: 5 * time.Second, // 重连间隔
|
||||
pingInterval: 2 * time.Second, // 发送间隔改为2秒
|
||||
curBackoff: initialBackoff, // 当前退避间隔
|
||||
pingInterval: 1 * time.Second, // 指标上报间隔(每秒采集)
|
||||
configInterval: 10 * time.Minute, // 配置上报间隔
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
@@ -147,10 +208,17 @@ func (w *WebSocketReporter) Start() {
|
||||
// Stop 停止WebSocket报告器
|
||||
func (w *WebSocketReporter) Stop() {
|
||||
w.cancel()
|
||||
w.connMutex.Lock()
|
||||
if w.conn != nil {
|
||||
w.conn.Close()
|
||||
}
|
||||
w.connMutex.Unlock()
|
||||
}
|
||||
|
||||
// backoffWithJitter 返回带随机抖动的退避时间(±25%)
|
||||
func backoffWithJitter(base time.Duration) time.Duration {
|
||||
jitter := time.Duration(float64(base) * (0.75 + rand.Float64()*0.5))
|
||||
return jitter
|
||||
}
|
||||
|
||||
// run 主运行循环
|
||||
@@ -167,23 +235,32 @@ func (w *WebSocketReporter) run() {
|
||||
|
||||
if needConnect {
|
||||
if err := w.connect(); err != nil {
|
||||
fmt.Printf("❌ WebSocket连接失败: %v,%v后重试\n", err, w.reconnectTime)
|
||||
wait := backoffWithJitter(w.curBackoff)
|
||||
fmt.Printf("❌ WebSocket连接失败: %v,%v后重试\n", err, wait)
|
||||
// 指数退避:翻倍当前退避间隔,上限 maxBackoff
|
||||
w.curBackoff *= 2
|
||||
if w.curBackoff > maxBackoff {
|
||||
w.curBackoff = maxBackoff
|
||||
}
|
||||
select {
|
||||
case <-time.After(w.reconnectTime):
|
||||
case <-time.After(wait):
|
||||
continue
|
||||
case <-w.ctx.Done():
|
||||
return
|
||||
}
|
||||
}
|
||||
// 连接成功:重置退避
|
||||
w.curBackoff = initialBackoff
|
||||
}
|
||||
|
||||
// 连接成功,开始发送消息
|
||||
if w.connected {
|
||||
w.handleConnection()
|
||||
} else {
|
||||
wait := backoffWithJitter(w.curBackoff)
|
||||
// 如果连接失败,等待重试
|
||||
select {
|
||||
case <-time.After(w.reconnectTime):
|
||||
case <-time.After(wait):
|
||||
continue
|
||||
case <-w.ctx.Done():
|
||||
return
|
||||
@@ -223,21 +300,14 @@ func (w *WebSocketReporter) connect() error {
|
||||
json.Unmarshal(b, &cfg)
|
||||
}
|
||||
|
||||
// 使用最新的配置重新构建 URL
|
||||
currentURL := "ws://" + w.addr + "/system-info?type=1&secret=" + w.secret + "&version=" + w.version +
|
||||
"&http=" + strconv.Itoa(cfg.Http) + "&tls=" + strconv.Itoa(cfg.Tls) + "&socks=" + strconv.Itoa(cfg.Socks)
|
||||
|
||||
u, err := url.Parse(currentURL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("解析URL失败: %v", err)
|
||||
}
|
||||
candidates := buildWebSocketCandidates(w.addr, w.secret, w.version, cfg.Http, cfg.Tls, cfg.Socks, w.preferredWSScheme)
|
||||
|
||||
dialer := websocket.DefaultDialer
|
||||
dialer.HandshakeTimeout = 10 * time.Second
|
||||
|
||||
conn, _, err := dialer.Dial(u.String(), nil)
|
||||
conn, usedURL, err := dialWebSocketWithFallback(dialer, candidates)
|
||||
if err != nil {
|
||||
return fmt.Errorf("连接WebSocket失败: %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// 如果在连接过程中已经有连接了,关闭新连接
|
||||
@@ -248,6 +318,9 @@ func (w *WebSocketReporter) connect() error {
|
||||
|
||||
w.conn = conn
|
||||
w.connected = true
|
||||
if scheme := detectWebSocketScheme(usedURL); scheme != "" {
|
||||
w.preferredWSScheme = scheme
|
||||
}
|
||||
_ = conn.SetReadDeadline(time.Now().Add(reporterReadWait))
|
||||
conn.SetPingHandler(func(appData string) error {
|
||||
_ = conn.SetReadDeadline(time.Now().Add(reporterReadWait))
|
||||
@@ -265,10 +338,145 @@ func (w *WebSocketReporter) connect() error {
|
||||
return nil
|
||||
})
|
||||
|
||||
fmt.Printf("✅ WebSocket连接建立成功 (http=%d, tls=%d, socks=%d)\n", cfg.Http, cfg.Tls, cfg.Socks)
|
||||
fmt.Printf("✅ WebSocket连接建立成功 (%s, http=%d, tls=%d, socks=%d)\n", sanitizeWebSocketURL(usedURL), cfg.Http, cfg.Tls, cfg.Socks)
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildWebSocketCandidates(addr string, secret string, version string, http int, tls int, socks int, preferredScheme string) []string {
|
||||
normalizedAddr, explicitScheme := normalizeReporterAddress(addr)
|
||||
if normalizedAddr == "" {
|
||||
normalizedAddr = strings.TrimSpace(addr)
|
||||
}
|
||||
|
||||
query := "/system-info?type=1&secret=" + url.QueryEscape(secret) + "&version=" + url.QueryEscape(version) +
|
||||
"&http=" + strconv.Itoa(http) + "&tls=" + strconv.Itoa(tls) + "&socks=" + strconv.Itoa(socks)
|
||||
|
||||
schemes := []string{"wss", "ws"}
|
||||
if mappedScheme := mapToWebSocketScheme(explicitScheme); mappedScheme != "" {
|
||||
if mappedScheme == "ws" {
|
||||
schemes = []string{"ws", "wss"}
|
||||
}
|
||||
} else if preferredScheme == "ws" {
|
||||
schemes = []string{"ws", "wss"}
|
||||
}
|
||||
|
||||
return []string{
|
||||
schemes[0] + "://" + normalizedAddr + query,
|
||||
schemes[1] + "://" + normalizedAddr + query,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeReporterAddress(addr string) (string, string) {
|
||||
raw := strings.TrimSpace(addr)
|
||||
if raw == "" {
|
||||
return "", ""
|
||||
}
|
||||
|
||||
scheme := ""
|
||||
if idx := strings.Index(raw, "://"); idx > 0 {
|
||||
scheme = strings.ToLower(strings.TrimSpace(raw[:idx]))
|
||||
if parsed, err := url.Parse(raw); err == nil {
|
||||
if host := strings.TrimSpace(parsed.Host); host != "" {
|
||||
return host, scheme
|
||||
}
|
||||
}
|
||||
raw = raw[idx+3:]
|
||||
}
|
||||
|
||||
if idx := strings.IndexAny(raw, "/?#"); idx >= 0 {
|
||||
raw = raw[:idx]
|
||||
}
|
||||
return strings.TrimSpace(raw), scheme
|
||||
}
|
||||
|
||||
func mapToWebSocketScheme(scheme string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(scheme)) {
|
||||
case "wss", "https":
|
||||
return "wss"
|
||||
case "ws", "http":
|
||||
return "ws"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func detectWebSocketScheme(rawURL string) string {
|
||||
if strings.HasPrefix(rawURL, "wss://") {
|
||||
return "wss"
|
||||
}
|
||||
if strings.HasPrefix(rawURL, "ws://") {
|
||||
return "ws"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func dialWebSocketWithFallback(dialer *websocket.Dialer, candidates []string) (*websocket.Conn, string, error) {
|
||||
if len(candidates) == 0 {
|
||||
return nil, "", fmt.Errorf("WebSocket候选地址为空")
|
||||
}
|
||||
|
||||
var errs []string
|
||||
for i, targetURL := range candidates {
|
||||
conn, resp, err := wsDial(dialer, targetURL)
|
||||
if err == nil {
|
||||
if i > 0 {
|
||||
fmt.Printf("↪️ WebSocket已自动回退成功: %s\n", sanitizeWebSocketURL(targetURL))
|
||||
}
|
||||
return conn, targetURL, nil
|
||||
}
|
||||
errMsg := formatWebSocketDialError(err, resp)
|
||||
errs = append(errs, fmt.Sprintf("%s => %s", sanitizeWebSocketURL(targetURL), errMsg))
|
||||
if i < len(candidates)-1 {
|
||||
fmt.Printf(
|
||||
"⚠️ WebSocket连接失败,准备从 %s 回退到 %s: %s\n",
|
||||
strings.ToUpper(detectWebSocketScheme(targetURL)),
|
||||
strings.ToUpper(detectWebSocketScheme(candidates[i+1])),
|
||||
errMsg,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return nil, "", fmt.Errorf("连接WebSocket失败(已尝试%d种协议): %s", len(candidates), strings.Join(errs, " | "))
|
||||
}
|
||||
|
||||
func sanitizeWebSocketURL(rawURL string) string {
|
||||
u, err := url.Parse(rawURL)
|
||||
if err != nil {
|
||||
return rawURL
|
||||
}
|
||||
|
||||
q := u.Query()
|
||||
if q.Get("secret") != "" {
|
||||
q.Set("secret", "***")
|
||||
u.RawQuery = q.Encode()
|
||||
}
|
||||
return u.String()
|
||||
}
|
||||
|
||||
func formatWebSocketDialError(err error, resp *http.Response) string {
|
||||
if err == nil {
|
||||
return ""
|
||||
}
|
||||
if resp == nil {
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("%s (HTTP %s)", err, resp.Status)
|
||||
if resp.Body == nil {
|
||||
return msg
|
||||
}
|
||||
|
||||
body, readErr := io.ReadAll(io.LimitReader(resp.Body, 256))
|
||||
if readErr != nil {
|
||||
return msg
|
||||
}
|
||||
bodyText := strings.TrimSpace(string(body))
|
||||
if bodyText == "" {
|
||||
return msg
|
||||
}
|
||||
return fmt.Sprintf("%s, body=%q", msg, bodyText)
|
||||
}
|
||||
|
||||
// handleConnection 处理WebSocket连接
|
||||
func (w *WebSocketReporter) handleConnection() {
|
||||
defer func() {
|
||||
@@ -285,15 +493,34 @@ func (w *WebSocketReporter) handleConnection() {
|
||||
// 启动消息接收goroutine
|
||||
go w.receiveMessages()
|
||||
|
||||
// 主发送循环
|
||||
ticker := time.NewTicker(w.pingInterval)
|
||||
defer ticker.Stop()
|
||||
// 指标上报 ticker
|
||||
metricTicker := time.NewTicker(w.pingInterval)
|
||||
defer metricTicker.Stop()
|
||||
|
||||
// 独立 WebSocket keepalive ping ticker
|
||||
pingTicker := time.NewTicker(wsPingInterval)
|
||||
defer pingTicker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-w.ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
|
||||
case <-pingTicker.C:
|
||||
// 发送 WebSocket ping 保活,独立于指标上报
|
||||
w.connMutex.Lock()
|
||||
conn := w.conn
|
||||
isConnected := w.connected
|
||||
w.connMutex.Unlock()
|
||||
if !isConnected || conn == nil {
|
||||
return
|
||||
}
|
||||
if err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(reporterWriteWait)); err != nil {
|
||||
fmt.Printf("❌ 发送WebSocket ping失败: %v,准备重连\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
case <-metricTicker.C:
|
||||
// 检查连接状态
|
||||
w.connMutex.Lock()
|
||||
isConnected := w.connected
|
||||
@@ -313,11 +540,35 @@ func (w *WebSocketReporter) handleConnection() {
|
||||
}
|
||||
}
|
||||
|
||||
var lastNetBytesReceived uint64
|
||||
var lastNetBytesTransmitted uint64
|
||||
var lastNetTime int64
|
||||
|
||||
var connInfoCached ConnectionInfo
|
||||
var connInfoCachedAt int64
|
||||
var connInfoCachedMu sync.Mutex
|
||||
|
||||
// collectSystemInfo 收集系统信息
|
||||
func (w *WebSocketReporter) collectSystemInfo() SystemInfo {
|
||||
networkStats := getNetworkStats()
|
||||
cpuInfo := getCPUInfo()
|
||||
memoryInfo := getMemoryInfo()
|
||||
diskInfo := getDiskInfo()
|
||||
loadInfo := getLoadInfo()
|
||||
connInfo := getConnectionInfo()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
var netInSpeed, netOutSpeed int64
|
||||
if lastNetTime > 0 {
|
||||
deltaMs := now - lastNetTime
|
||||
if deltaMs > 0 {
|
||||
netInSpeed = int64(float64(networkStats.BytesRecvDelta) * 1000 / float64(deltaMs))
|
||||
netOutSpeed = int64(float64(networkStats.BytesSentDelta) * 1000 / float64(deltaMs))
|
||||
}
|
||||
}
|
||||
lastNetBytesReceived = networkStats.BytesReceived
|
||||
lastNetBytesTransmitted = networkStats.BytesTransmitted
|
||||
lastNetTime = now
|
||||
|
||||
return SystemInfo{
|
||||
Uptime: getUptime(),
|
||||
@@ -325,9 +576,48 @@ func (w *WebSocketReporter) collectSystemInfo() SystemInfo {
|
||||
BytesTransmitted: networkStats.BytesTransmitted,
|
||||
CPUUsage: cpuInfo.Usage,
|
||||
MemoryUsage: memoryInfo.Usage,
|
||||
DiskUsage: diskInfo.Usage,
|
||||
Load1: loadInfo.Load1,
|
||||
Load5: loadInfo.Load5,
|
||||
Load15: loadInfo.Load15,
|
||||
TCPConns: connInfo.TCPConns,
|
||||
UDPConns: connInfo.UDPConns,
|
||||
NetInSpeed: netInSpeed,
|
||||
NetOutSpeed: netOutSpeed,
|
||||
}
|
||||
}
|
||||
|
||||
// encryptPayload 加密 JSON 数据,返回加密后的消息字节(若加密失败则回退到原始数据)
|
||||
func (w *WebSocketReporter) encryptPayload(jsonData []byte) []byte {
|
||||
if w.aesCrypto == nil {
|
||||
return jsonData
|
||||
}
|
||||
|
||||
encryptedData, err := w.aesCrypto.Encrypt(jsonData)
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ 加密失败,发送原始数据: %v\n", err)
|
||||
return jsonData
|
||||
}
|
||||
|
||||
encryptedMessage := map[string]interface{}{
|
||||
"encrypted": true,
|
||||
"data": encryptedData,
|
||||
"timestamp": time.Now().Unix(),
|
||||
}
|
||||
messageData, err := json.Marshal(encryptedMessage)
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ 序列化加密消息失败,发送原始数据: %v\n", err)
|
||||
return jsonData
|
||||
}
|
||||
return messageData
|
||||
}
|
||||
|
||||
// metricEnvelope wraps SystemInfo with a type field for fast identification on the panel side.
|
||||
type metricEnvelope struct {
|
||||
Type string `json:"type"`
|
||||
Data SystemInfo `json:"data"`
|
||||
}
|
||||
|
||||
// sendSystemInfo 发送系统信息
|
||||
func (w *WebSocketReporter) sendSystemInfo(sysInfo SystemInfo) error {
|
||||
w.connMutex.Lock()
|
||||
@@ -337,42 +627,19 @@ func (w *WebSocketReporter) sendSystemInfo(sysInfo SystemInfo) error {
|
||||
return fmt.Errorf("连接未建立")
|
||||
}
|
||||
|
||||
// 转换为JSON
|
||||
jsonData, err := json.Marshal(sysInfo)
|
||||
// 使用 type:"metric" 信封包装,Panel 可通过 type 字段直接识别指标消息
|
||||
envelope := metricEnvelope{Type: "metric", Data: sysInfo}
|
||||
jsonData, err := json.Marshal(envelope)
|
||||
if err != nil {
|
||||
return fmt.Errorf("序列化系统信息失败: %v", err)
|
||||
}
|
||||
|
||||
var messageData []byte
|
||||
messageData := w.encryptPayload(jsonData)
|
||||
|
||||
// 如果有加密器,则加密数据
|
||||
if w.aesCrypto != nil {
|
||||
encryptedData, err := w.aesCrypto.Encrypt(jsonData)
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ 加密失败,发送原始数据: %v\n", err)
|
||||
messageData = jsonData
|
||||
} else {
|
||||
// 创建加密消息包装器
|
||||
encryptedMessage := map[string]interface{}{
|
||||
"encrypted": true,
|
||||
"data": encryptedData,
|
||||
"timestamp": time.Now().Unix(),
|
||||
}
|
||||
messageData, err = json.Marshal(encryptedMessage)
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ 序列化加密消息失败,发送原始数据: %v\n", err)
|
||||
messageData = jsonData
|
||||
}
|
||||
}
|
||||
} else {
|
||||
messageData = jsonData
|
||||
}
|
||||
|
||||
// 设置写入超时
|
||||
w.conn.SetWriteDeadline(time.Now().Add(5 * time.Second))
|
||||
|
||||
if err := w.conn.WriteMessage(websocket.TextMessage, messageData); err != nil {
|
||||
w.connected = false // 标记连接已断开
|
||||
w.connected = false
|
||||
return fmt.Errorf("写入消息失败: %v", err)
|
||||
}
|
||||
|
||||
@@ -381,23 +648,19 @@ func (w *WebSocketReporter) sendSystemInfo(sysInfo SystemInfo) error {
|
||||
|
||||
// receiveMessages 接收服务端发送的消息
|
||||
func (w *WebSocketReporter) receiveMessages() {
|
||||
// 获取连接引用一次即可,连接生命周期由 handleConnection 管理
|
||||
w.connMutex.Lock()
|
||||
conn := w.conn
|
||||
w.connMutex.Unlock()
|
||||
if conn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-w.ctx.Done():
|
||||
return
|
||||
default:
|
||||
w.connMutex.Lock()
|
||||
conn := w.conn
|
||||
connected := w.connected
|
||||
w.connMutex.Unlock()
|
||||
|
||||
if conn == nil || !connected {
|
||||
return
|
||||
}
|
||||
|
||||
// 设置读取超时
|
||||
conn.SetReadDeadline(time.Now().Add(reporterReadWait))
|
||||
|
||||
messageType, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
|
||||
@@ -485,12 +748,8 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
||||
}
|
||||
|
||||
if cmdMsg.Type != "call" {
|
||||
// 其他状态变更命令保持同步,确保顺序执行
|
||||
if cmdMsg.Type == "TcpPing" || cmdMsg.Type == "UpgradeAgent" || cmdMsg.Type == "RollbackAgent" {
|
||||
go w.routeCommand(cmdMsg)
|
||||
} else {
|
||||
w.routeCommand(cmdMsg)
|
||||
}
|
||||
// 所有命令统一异步执行,避免阻塞消息接收循环
|
||||
go w.routeCommand(cmdMsg)
|
||||
}
|
||||
} else {
|
||||
// 处理普通消息
|
||||
@@ -501,12 +760,8 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
||||
return
|
||||
}
|
||||
if cmdMsg.Type != "call" {
|
||||
// 其他状态变更命令保持同步,确保顺序执行
|
||||
if cmdMsg.Type == "TcpPing" || cmdMsg.Type == "UpgradeAgent" || cmdMsg.Type == "RollbackAgent" {
|
||||
go w.routeCommand(cmdMsg)
|
||||
} else {
|
||||
w.routeCommand(cmdMsg)
|
||||
}
|
||||
// 所有命令统一异步执行,避免阻塞消息接收循环
|
||||
go w.routeCommand(cmdMsg)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -590,6 +845,13 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
response.Data = tcpPingResult
|
||||
// needSaveConfig = false (默认值)
|
||||
|
||||
// Service monitor check (read-only)
|
||||
case "ServiceMonitorCheck":
|
||||
var checkResult ServiceMonitorCheckResult
|
||||
checkResult, err = w.handleServiceMonitorCheck(cmd.Data)
|
||||
response.Type = "ServiceMonitorCheckResponse"
|
||||
response.Data = checkResult
|
||||
|
||||
// Protocol blocking switches
|
||||
case "SetProtocol":
|
||||
err = w.handleSetProtocol(cmd.Data)
|
||||
@@ -1173,30 +1435,7 @@ func (w *WebSocketReporter) sendResponse(response CommandResponse) {
|
||||
return
|
||||
}
|
||||
|
||||
var messageData []byte
|
||||
|
||||
// 如果有加密器,则加密数据
|
||||
if w.aesCrypto != nil {
|
||||
encryptedData, err := w.aesCrypto.Encrypt(jsonData)
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ 加密响应失败,发送原始数据: %v\n", err)
|
||||
messageData = jsonData
|
||||
} else {
|
||||
// 创建加密消息包装器
|
||||
encryptedMessage := map[string]interface{}{
|
||||
"encrypted": true,
|
||||
"data": encryptedData,
|
||||
"timestamp": time.Now().Unix(),
|
||||
}
|
||||
messageData, err = json.Marshal(encryptedMessage)
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ 序列化加密响应失败,发送原始数据: %v\n", err)
|
||||
messageData = jsonData
|
||||
}
|
||||
}
|
||||
} else {
|
||||
messageData = jsonData
|
||||
}
|
||||
messageData := w.encryptPayload(jsonData)
|
||||
|
||||
// 检查消息大小,如果超过10MB则记录警告
|
||||
if len(messageData) > 10*1024*1024 {
|
||||
@@ -1245,17 +1484,21 @@ func getNetworkStats() NetworkStats {
|
||||
return stats
|
||||
}
|
||||
|
||||
// 汇总所有非回环接口的流量
|
||||
for _, io := range ioCounters {
|
||||
// 跳过回环接口
|
||||
if io.Name == "lo" || strings.HasPrefix(io.Name, "lo") {
|
||||
continue
|
||||
}
|
||||
|
||||
stats.BytesReceived += io.BytesRecv
|
||||
stats.BytesTransmitted += io.BytesSent
|
||||
}
|
||||
|
||||
if lastNetBytesReceived > 0 && stats.BytesReceived >= lastNetBytesReceived {
|
||||
stats.BytesRecvDelta = stats.BytesReceived - lastNetBytesReceived
|
||||
}
|
||||
if lastNetBytesTransmitted > 0 && stats.BytesTransmitted >= lastNetBytesTransmitted {
|
||||
stats.BytesSentDelta = stats.BytesTransmitted - lastNetBytesTransmitted
|
||||
}
|
||||
|
||||
return stats
|
||||
}
|
||||
|
||||
@@ -1263,8 +1506,8 @@ func getNetworkStats() NetworkStats {
|
||||
func getCPUInfo() CPUInfo {
|
||||
var cpuInfo CPUInfo
|
||||
|
||||
// 获取CPU使用率
|
||||
percentages, err := cpu.Percent(time.Second, false)
|
||||
// 获取CPU使用率 (non-blocking)
|
||||
percentages, err := cpu.Percent(0, false)
|
||||
if err == nil && len(percentages) > 0 {
|
||||
cpuInfo.Usage = percentages[0]
|
||||
}
|
||||
@@ -1286,11 +1529,75 @@ func getMemoryInfo() MemoryInfo {
|
||||
return memInfo
|
||||
}
|
||||
|
||||
// getDiskInfo 获取磁盘信息
|
||||
func getDiskInfo() DiskInfo {
|
||||
var diskInfo DiskInfo
|
||||
|
||||
usage, err := disk.Usage("/")
|
||||
if err != nil {
|
||||
return diskInfo
|
||||
}
|
||||
|
||||
diskInfo.Usage = usage.UsedPercent
|
||||
|
||||
return diskInfo
|
||||
}
|
||||
|
||||
// getLoadInfo 获取负载信息
|
||||
func getLoadInfo() LoadInfo {
|
||||
var loadInfo LoadInfo
|
||||
|
||||
avg, err := load.Avg()
|
||||
if err != nil {
|
||||
return loadInfo
|
||||
}
|
||||
|
||||
loadInfo.Load1 = avg.Load1
|
||||
loadInfo.Load5 = avg.Load5
|
||||
loadInfo.Load15 = avg.Load15
|
||||
|
||||
return loadInfo
|
||||
}
|
||||
|
||||
// getConnectionInfo 获取连接信息
|
||||
func getConnectionInfo() ConnectionInfo {
|
||||
now := time.Now().UnixMilli()
|
||||
const refreshEveryMs = int64((15 * time.Second) / time.Millisecond)
|
||||
|
||||
connInfoCachedMu.Lock()
|
||||
if connInfoCachedAt > 0 && now-connInfoCachedAt < refreshEveryMs {
|
||||
v := connInfoCached
|
||||
connInfoCachedMu.Unlock()
|
||||
return v
|
||||
}
|
||||
connInfoCachedMu.Unlock()
|
||||
|
||||
var connInfo ConnectionInfo
|
||||
|
||||
connStats, err := psnet.Connections("tcp")
|
||||
if err == nil {
|
||||
connInfo.TCPConns = int64(len(connStats))
|
||||
}
|
||||
|
||||
udpStats, err := psnet.Connections("udp")
|
||||
if err == nil {
|
||||
connInfo.UDPConns = int64(len(udpStats))
|
||||
}
|
||||
|
||||
connInfoCachedMu.Lock()
|
||||
connInfoCached = connInfo
|
||||
connInfoCachedAt = now
|
||||
connInfoCachedMu.Unlock()
|
||||
|
||||
return connInfo
|
||||
}
|
||||
|
||||
// StartWebSocketReporterWithConfig 使用配置字段启动WebSocket报告器
|
||||
func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls int, socks int, version string) *WebSocketReporter {
|
||||
|
||||
// 构建初始 WebSocket URL
|
||||
fullURL := "ws://" + addr + "/system-info?type=1&secret=" + secret + "&version=" + version + "&http=" + strconv.Itoa(http) + "&tls=" + strconv.Itoa(tls) + "&socks=" + strconv.Itoa(socks)
|
||||
candidates := buildWebSocketCandidates(addr, secret, version, http, tls, socks, "")
|
||||
fullURL := candidates[0]
|
||||
|
||||
fmt.Printf("🔗 WebSocket连接URL: %s\n", fullURL)
|
||||
|
||||
@@ -1366,6 +1673,210 @@ func (w *WebSocketReporter) handleTcpPing(data interface{}) (TcpPingResponse, er
|
||||
return response, nil
|
||||
}
|
||||
|
||||
// handleServiceMonitorCheck executes a service monitor check on this node.
|
||||
// It always returns a result (command execution is considered successful even if the check fails).
|
||||
func (w *WebSocketReporter) handleServiceMonitorCheck(data interface{}) (ServiceMonitorCheckResult, error) {
|
||||
jsonData, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return ServiceMonitorCheckResult{}, fmt.Errorf("序列化检查数据失败: %v", err)
|
||||
}
|
||||
|
||||
var req ServiceMonitorCheckRequest
|
||||
if err := json.Unmarshal(jsonData, &req); err != nil {
|
||||
return ServiceMonitorCheckResult{}, fmt.Errorf("解析检查请求失败: %v", err)
|
||||
}
|
||||
|
||||
checkType := strings.ToLower(strings.TrimSpace(req.Type))
|
||||
target := strings.TrimSpace(req.Target)
|
||||
res := ServiceMonitorCheckResult{MonitorID: req.MonitorID}
|
||||
|
||||
if checkType != "tcp" && checkType != "icmp" {
|
||||
res.Success = false
|
||||
res.ErrorMessage = "不支持的检查类型"
|
||||
return res, nil
|
||||
}
|
||||
if target == "" {
|
||||
res.Success = false
|
||||
res.ErrorMessage = "检查目标为空"
|
||||
return res, nil
|
||||
}
|
||||
|
||||
timeoutSec := req.TimeoutSec
|
||||
if timeoutSec <= 0 {
|
||||
timeoutSec = 5
|
||||
}
|
||||
timeout := time.Duration(timeoutSec) * time.Second
|
||||
|
||||
start := time.Now()
|
||||
|
||||
switch checkType {
|
||||
case "tcp":
|
||||
// Validate and normalize host:port.
|
||||
_, _, splitErr := net.SplitHostPort(target)
|
||||
if splitErr != nil {
|
||||
res.Success = false
|
||||
res.ErrorMessage = "无效的TCP目标"
|
||||
res.LatencyMs = float64(time.Since(start).Milliseconds())
|
||||
return res, nil
|
||||
}
|
||||
conn, dialErr := net.DialTimeout("tcp", target, timeout)
|
||||
res.LatencyMs = float64(time.Since(start).Milliseconds())
|
||||
if dialErr != nil {
|
||||
res.Success = false
|
||||
res.ErrorMessage = dialErr.Error()
|
||||
return res, nil
|
||||
}
|
||||
_ = conn.Close()
|
||||
res.Success = true
|
||||
return res, nil
|
||||
|
||||
case "icmp":
|
||||
rtt, pingErr := icmpPing(target, timeout)
|
||||
res.LatencyMs = float64(rtt.Milliseconds())
|
||||
if pingErr != nil {
|
||||
res.Success = false
|
||||
res.ErrorMessage = pingErr.Error()
|
||||
return res, nil
|
||||
}
|
||||
res.Success = true
|
||||
return res, nil
|
||||
}
|
||||
|
||||
res.Success = false
|
||||
res.ErrorMessage = "未知错误"
|
||||
res.LatencyMs = float64(time.Since(start).Milliseconds())
|
||||
return res, nil
|
||||
}
|
||||
|
||||
func icmpPing(target string, timeout time.Duration) (time.Duration, error) {
|
||||
start := time.Now()
|
||||
|
||||
target = strings.TrimSpace(target)
|
||||
if target == "" {
|
||||
return time.Since(start), fmt.Errorf("无效的ICMP目标")
|
||||
}
|
||||
// Avoid accepting URL-like targets.
|
||||
if strings.Contains(target, "://") {
|
||||
return time.Since(start), fmt.Errorf("无效的ICMP目标")
|
||||
}
|
||||
if strings.HasPrefix(target, "[") && strings.HasSuffix(target, "]") {
|
||||
target = strings.TrimSuffix(strings.TrimPrefix(target, "["), "]")
|
||||
}
|
||||
|
||||
ipAddr, err := net.ResolveIPAddr("ip", target)
|
||||
if err != nil || ipAddr == nil || ipAddr.IP == nil {
|
||||
if err == nil {
|
||||
err = fmt.Errorf("unknown address")
|
||||
}
|
||||
return time.Since(start), fmt.Errorf("解析目标失败: %v", err)
|
||||
}
|
||||
|
||||
isV4 := ipAddr.IP.To4() != nil
|
||||
listenAddr := "0.0.0.0"
|
||||
proto := 1
|
||||
var echoType icmp.Type = ipv4.ICMPTypeEcho
|
||||
var echoReplyType icmp.Type = ipv4.ICMPTypeEchoReply
|
||||
networks := []string{"udp4", "ip4:icmp"}
|
||||
if !isV4 {
|
||||
listenAddr = "::"
|
||||
proto = 58
|
||||
echoType = ipv6.ICMPTypeEchoRequest
|
||||
echoReplyType = ipv6.ICMPTypeEchoReply
|
||||
networks = []string{"udp6", "ip6:ipv6-icmp"}
|
||||
}
|
||||
|
||||
var conn *icmp.PacketConn
|
||||
selectedNetwork := ""
|
||||
var lastErr error
|
||||
for _, nw := range networks {
|
||||
c, err := icmp.ListenPacket(nw, listenAddr)
|
||||
if err == nil {
|
||||
conn = c
|
||||
selectedNetwork = nw
|
||||
break
|
||||
}
|
||||
lastErr = err
|
||||
}
|
||||
if conn == nil {
|
||||
if lastErr != nil {
|
||||
return time.Since(start), fmt.Errorf("创建ICMP连接失败: %v", lastErr)
|
||||
}
|
||||
return time.Since(start), fmt.Errorf("创建ICMP连接失败")
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
id := os.Getpid() & 0xffff
|
||||
seq := 1
|
||||
|
||||
wm := icmp.Message{
|
||||
Type: echoType,
|
||||
Code: 0,
|
||||
Body: &icmp.Echo{
|
||||
ID: id,
|
||||
Seq: seq,
|
||||
Data: []byte("FLVX-PING"),
|
||||
},
|
||||
}
|
||||
wb, err := wm.Marshal(nil)
|
||||
if err != nil {
|
||||
return time.Since(start), err
|
||||
}
|
||||
|
||||
_ = conn.SetDeadline(time.Now().Add(timeout))
|
||||
|
||||
var dst net.Addr
|
||||
if strings.HasPrefix(selectedNetwork, "udp") {
|
||||
dst = &net.UDPAddr{IP: ipAddr.IP, Zone: ipAddr.Zone}
|
||||
} else {
|
||||
dst = &net.IPAddr{IP: ipAddr.IP, Zone: ipAddr.Zone}
|
||||
}
|
||||
|
||||
if _, err := conn.WriteTo(wb, dst); err != nil {
|
||||
return time.Since(start), err
|
||||
}
|
||||
|
||||
addrIP := func(a net.Addr) net.IP {
|
||||
switch v := a.(type) {
|
||||
case *net.IPAddr:
|
||||
return v.IP
|
||||
case *net.UDPAddr:
|
||||
return v.IP
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
rb := make([]byte, 1500)
|
||||
for {
|
||||
n, peer, err := conn.ReadFrom(rb)
|
||||
if err != nil {
|
||||
return time.Since(start), err
|
||||
}
|
||||
if p := addrIP(peer); p != nil && !p.Equal(ipAddr.IP) {
|
||||
continue
|
||||
}
|
||||
rm, err := icmp.ParseMessage(proto, rb[:n])
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if rm.Type != echoReplyType {
|
||||
continue
|
||||
}
|
||||
echo, ok := rm.Body.(*icmp.Echo)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if echo.Seq != seq {
|
||||
continue
|
||||
}
|
||||
// For non-privileged endpoints, the kernel may choose the ID.
|
||||
if !strings.HasPrefix(selectedNetwork, "udp") && echo.ID != id {
|
||||
continue
|
||||
}
|
||||
return time.Since(start), nil
|
||||
}
|
||||
}
|
||||
|
||||
// tcpPingHost 执行TCP连接测试,返回平均连接时间和失败率
|
||||
func tcpPingHost(ip string, port int, count int, timeoutMs int) (float64, float64, error) {
|
||||
var totalTime float64
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func TestBuildWebSocketCandidatesSecureFirst(t *testing.T) {
|
||||
candidates := buildWebSocketCandidates("panel.example.com:443", "abc", "2.0.2", 1, 0, 1, "")
|
||||
|
||||
if len(candidates) != 2 {
|
||||
t.Fatalf("expected 2 candidates, got %d", len(candidates))
|
||||
}
|
||||
if !strings.HasPrefix(candidates[0], "wss://") {
|
||||
t.Fatalf("expected first candidate to start with wss://, got %s", candidates[0])
|
||||
}
|
||||
if !strings.HasPrefix(candidates[1], "ws://") {
|
||||
t.Fatalf("expected second candidate to start with ws://, got %s", candidates[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildWebSocketCandidatesUsesPreferredScheme(t *testing.T) {
|
||||
candidates := buildWebSocketCandidates("panel.example.com:443", "abc", "2.0.2", 1, 0, 1, "ws")
|
||||
|
||||
if len(candidates) != 2 {
|
||||
t.Fatalf("expected 2 candidates, got %d", len(candidates))
|
||||
}
|
||||
if !strings.HasPrefix(candidates[0], "ws://") {
|
||||
t.Fatalf("expected preferred ws:// candidate first, got %s", candidates[0])
|
||||
}
|
||||
if !strings.HasPrefix(candidates[1], "wss://") {
|
||||
t.Fatalf("expected fallback wss:// candidate second, got %s", candidates[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildWebSocketCandidatesNormalizesSchemePrefixedAddr(t *testing.T) {
|
||||
candidates := buildWebSocketCandidates("https://panel.example.com:443/path?q=1", "abc", "2.0.2", 0, 0, 0, "")
|
||||
|
||||
if len(candidates) != 2 {
|
||||
t.Fatalf("expected 2 candidates, got %d", len(candidates))
|
||||
}
|
||||
if !strings.HasPrefix(candidates[0], "wss://panel.example.com:443/") {
|
||||
t.Fatalf("expected normalized wss candidate, got %s", candidates[0])
|
||||
}
|
||||
if !strings.HasPrefix(candidates[1], "ws://panel.example.com:443/") {
|
||||
t.Fatalf("expected normalized ws fallback candidate, got %s", candidates[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestDialWebSocketWithFallbackTriesWSAfterWSSFailure(t *testing.T) {
|
||||
orig := wsDial
|
||||
defer func() { wsDial = orig }()
|
||||
|
||||
var attempts []string
|
||||
wsDial = func(_ *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||
attempts = append(attempts, rawURL)
|
||||
if strings.HasPrefix(rawURL, "wss://") {
|
||||
return nil, nil, errors.New("tls failed")
|
||||
}
|
||||
return &websocket.Conn{}, nil, nil
|
||||
}
|
||||
|
||||
_, usedURL, err := dialWebSocketWithFallback(
|
||||
&websocket.Dialer{},
|
||||
[]string{
|
||||
"wss://panel.example.com/system-info?type=1&secret=abc",
|
||||
"ws://panel.example.com/system-info?type=1&secret=abc",
|
||||
},
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("expected fallback success, got err=%v", err)
|
||||
}
|
||||
if !strings.HasPrefix(usedURL, "ws://") {
|
||||
t.Fatalf("expected fallback ws:// url, got %s", usedURL)
|
||||
}
|
||||
if len(attempts) != 2 {
|
||||
t.Fatalf("expected 2 attempts, got %d", len(attempts))
|
||||
}
|
||||
if !strings.HasPrefix(attempts[0], "wss://") || !strings.HasPrefix(attempts[1], "ws://") {
|
||||
t.Fatalf("unexpected attempt order: %#v", attempts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectWebSocketScheme(t *testing.T) {
|
||||
if detectWebSocketScheme("wss://panel.example.com/system-info") != "wss" {
|
||||
t.Fatalf("expected wss detection")
|
||||
}
|
||||
if detectWebSocketScheme("ws://panel.example.com/system-info") != "ws" {
|
||||
t.Fatalf("expected ws detection")
|
||||
}
|
||||
if detectWebSocketScheme("http://panel.example.com/system-info") != "" {
|
||||
t.Fatalf("expected empty detection for non-websocket scheme")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeWebSocketURL(t *testing.T) {
|
||||
raw := "wss://panel.example.com/system-info?type=1&secret=abc&version=2.0.2"
|
||||
sanitized := sanitizeWebSocketURL(raw)
|
||||
|
||||
if strings.Contains(sanitized, "secret=abc") {
|
||||
t.Fatalf("expected secret to be masked, got %s", sanitized)
|
||||
}
|
||||
if !strings.Contains(sanitized, "secret=%2A%2A%2A") {
|
||||
t.Fatalf("expected masked secret in url, got %s", sanitized)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
|
||||
err := errors.New("websocket: bad handshake")
|
||||
resp := &http.Response{
|
||||
Status: "403 Forbidden",
|
||||
Body: io.NopCloser(strings.NewReader("forbidden")),
|
||||
}
|
||||
|
||||
msg := formatWebSocketDialError(err, resp)
|
||||
if !strings.Contains(msg, "HTTP 403 Forbidden") {
|
||||
t.Fatalf("expected status in message, got %s", msg)
|
||||
}
|
||||
if !strings.Contains(msg, "forbidden") {
|
||||
t.Fatalf("expected response body in message, got %s", msg)
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,8 @@ nav:
|
||||
- 首页: index.md
|
||||
- 安装部署: install.md
|
||||
- 使用指南: usage.md
|
||||
- AI Skill 接入: ai-skill.md
|
||||
- PostgreSQL: postgresql.md
|
||||
- 常见问题: faq.md
|
||||
|
||||
markdown_extensions:
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
schema: spec-driven
|
||||
created: 2026-02-17
|
||||
@@ -1,29 +0,0 @@
|
||||
## Context
|
||||
|
||||
FLVX is a distributed system consisting of a central management panel (Backend + Frontend) and multiple forwarding agents (Nodes). The backend manages configuration, users, and billing, while agents handle the actual traffic forwarding using a modified GOST v3 stack. Communication between the panel and agents is secured and synchronized.
|
||||
|
||||
## Goals / Non-Goals
|
||||
|
||||
**Goals:**
|
||||
- Document the high-level architecture of the system.
|
||||
- Describe the data model for users, tunnels, and nodes.
|
||||
- Explain the communication protocol between Panel and Agent.
|
||||
- Detail the authentication and authorization mechanisms.
|
||||
|
||||
**Non-Goals:**
|
||||
- Refactoring the existing architecture.
|
||||
- Detailed code-level documentation of every function.
|
||||
- Changing the database schema.
|
||||
|
||||
## Decisions
|
||||
|
||||
- **Architecture**: The system follows a client-server model where the Panel acts as the server and Agents act as clients that pull configuration and push status.
|
||||
- **Data Model**: Core entities are Users, Nodes (Agents), Tunnels (Groups of rules), and Forwarding Rules.
|
||||
- **Communication**: Agents use a heartbeat mechanism to report status and fetch configuration updates. The protocol uses AES encryption with a pre-shared key (Node Secret).
|
||||
- **Authentication**: JWT for Frontend-Backend communication; API Key (Node Secret) for Agent-Backend communication.
|
||||
|
||||
## Risks / Trade-offs
|
||||
|
||||
- **Security**: The security of the agent communication relies heavily on the secrecy of the Node Secret.
|
||||
- **Scalability**: Centralized management might become a bottleneck with a very large number of agents.
|
||||
- **Complexity**: Synchronizing state across distributed agents introduces complexity in handling failures and inconsistencies.
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user