mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| feb357ff17 | |||
| 34581e0d18 | |||
| 61c5b5e759 | |||
| c8eb780c67 | |||
| 4bdfa50b0c |
@@ -1,12 +1,12 @@
|
||||
# PROJECT KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Thu Feb 19 2026
|
||||
**Commit:** 137c34e
|
||||
**Generated:** Thu Feb 26 2026
|
||||
**Commit:** 21008cc
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.4-rc2
|
||||
**Tag:** 2.1.5-rc15
|
||||
|
||||
## 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
|
||||
```
|
||||
@@ -14,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)
|
||||
│ └── 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
|
||||
@@ -29,12 +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) |
|
||||
| **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 |
|
||||
@@ -43,7 +48,9 @@ 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
|
||||
- **Auth**: `Authorization` header carries the raw JWT token (no `Bearer` prefix) between `vite-frontend/` and `go-backend/`.
|
||||
@@ -52,6 +59,7 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
- **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`.
|
||||
@@ -60,6 +68,8 @@ FLVX (formerly Flux Panel) is a traffic forwarding management system built on a
|
||||
- **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
|
||||
@@ -75,15 +85,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.
|
||||
@@ -91,7 +107,9 @@ 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.
|
||||
- PR `#144` (shadcn migration) and PR `#142` (user-group binding) are merged into `main`; release tag `2.1.4-rc2` points to commit `137c34e`.
|
||||
- Button visual parity relies on `vite-frontend/src/shadcn-bridge/heroui/button.tsx` color mapping + `vite-frontend/src/styles/tailwind-theme.pcss` token export.
|
||||
- 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.
|
||||
@@ -0,0 +1,148 @@
|
||||
# 限速功能重构实施计划
|
||||
|
||||
## 一、需求概述
|
||||
|
||||
**原始需求**: 限速功能当前绑定到具体隧道,需要改为不绑定隧道,创建限速后可以自由在隧道上限速,也可以在转发上限速。
|
||||
|
||||
**核心变更**:
|
||||
1. 限速规则(SpeedLimit)与隧道的绑定关系改为可选
|
||||
2. 转发(Forward)支持独立的限速规则
|
||||
|
||||
---
|
||||
|
||||
## 二、实施计划清单
|
||||
|
||||
### 2.0 计划状态(审计更新:2026-02-26)
|
||||
|
||||
- 总体状态:**进行中(未验收通过)**
|
||||
- 已完成:模型、仓储查询、限速 CRUD、控制面优先级、限速页与类型改造、编译与测试通过
|
||||
- 未完成:**Forward 独立限速写入链路**(前端表单 -> API handler -> repository 落库 `forward.speed_id`)
|
||||
|
||||
### 2.1 后端模型层 (Model)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| M1 | SpeedLimit.TunnelID 改为 sql.NullInt64 (可空) | `go-backend/internal/store/model/model.go` | ✅ 完成 |
|
||||
| M2 | SpeedLimit.TunnelName 改为 sql.NullString (可空) | `go-backend/internal/store/model/model.go` | ✅ 完成 |
|
||||
| M3 | Forward 添加 SpeedID sql.NullInt64 字段 | `go-backend/internal/store/model/model.go` | ✅ 完成 |
|
||||
| M4 | ForwardRecord 添加 SpeedID sql.NullInt64 字段 | `go-backend/internal/store/model/model.go` | ✅ 完成 |
|
||||
| M5 | SpeedLimitBackup.TunnelID 改为指针类型 | `go-backend/internal/store/model/model.go` | ✅ 完成 |
|
||||
| M6 | ForwardBackup 添加 SpeedID *int64 字段 | `go-backend/internal/store/model/model.go` | ✅ 完成 |
|
||||
|
||||
### 2.2 后端仓储层 (Repository)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| R1 | ListSpeedLimits() 返回可空 tunnelId/tunnelName | `go-backend/internal/store/repo/repository.go` | ✅ 完成 |
|
||||
| R2 | ListForwards() 返回 speedId 字段 | `go-backend/internal/store/repo/repository.go` | ✅ 完成 |
|
||||
| R3 | CreateSpeedLimit() 参数 tunnelID 改为 *int64 | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| R4 | UpdateSpeedLimit() 参数 tunnelID 改为 *int64 | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| R5 | GetSpeedLimitTunnelID() 返回 sql.NullInt64 | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| R6 | exportSpeedLimits() 处理可空字段 | `go-backend/internal/store/repo/repository.go` | ✅ 完成 |
|
||||
| R7 | importSpeedLimits() 处理可空字段 | `go-backend/internal/store/repo/repository.go` | ✅ 完成 |
|
||||
| R8 | GetSpeedLimitSpeed() 新增方法 | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 |
|
||||
| R9 | ListForwardsByTunnel() 返回 SpeedID | `go-backend/internal/store/repo/repository_control.go` | ✅ 完成 |
|
||||
| R10 | ListActiveForwardsByUser() 返回 SpeedID | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 |
|
||||
| R11 | ListActiveForwardsByUserTunnel() 返回 SpeedID | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 |
|
||||
| R12 | GetForwardRecord() 返回 SpeedID | `go-backend/internal/store/repo/repository_flow.go` | ✅ 完成 |
|
||||
|
||||
### 2.3 后端处理器层 (Handler)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| H1 | speedLimitCreate 处理可选 tunnelId | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
| H2 | speedLimitUpdate 处理可选 tunnelId | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
| H3 | speedLimitDelete 处理可空 tunnelID | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
|
||||
### 2.4 后端控制平面 (Control Plane)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| C1 | syncForwardServices 优先使用 Forward.SpeedID | `go-backend/internal/http/handler/control_plane.go` | ✅ 完成 |
|
||||
| C2 | 回退到 UserTunnel 的 speed limit | `go-backend/internal/http/handler/control_plane.go` | ✅ 完成 |
|
||||
|
||||
### 2.5 前端类型定义 (TypeScript Types)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| T1 | SpeedLimitApiItem.tunnelId 改为可选 | `vite-frontend/src/api/types.ts` | ✅ 完成 |
|
||||
| T2 | ForwardApiItem 添加 speedId 字段 | `vite-frontend/src/api/types.ts` | ✅ 完成 |
|
||||
| T3 | ForwardMutationPayload 添加 speedId 字段 | `vite-frontend/src/api/types.ts` | ✅ 完成 |
|
||||
| T4 | SpeedLimitMutationPayload.tunnelId 改为可选 | `vite-frontend/src/api/types.ts` | ✅ 完成 |
|
||||
|
||||
### 2.6 前端页面组件
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| F1 | SpeedLimitRule 接口更新 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 |
|
||||
| F2 | SpeedLimitForm 接口更新 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 |
|
||||
| F3 | validateForm 移除 tunnelId 必填校验 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 |
|
||||
| F4 | Select 组件改为可选 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 |
|
||||
| F5 | 显示"未绑定"状态 | `vite-frontend/src/pages/limit.tsx` | ✅ 完成 |
|
||||
|
||||
### 2.7 编译验证
|
||||
|
||||
| 序号 | 任务 | 状态 |
|
||||
|------|------|------|
|
||||
| B1 | Go 后端编译通过 | ✅ 完成 |
|
||||
| B2 | TypeScript 类型检查通过 | ✅ 完成 |
|
||||
| B3 | `go test ./...` 全量通过 | ✅ 完成 |
|
||||
| B4 | `go test ./tests/contract/... -run SpeedLimit` 通过 | ✅ 完成 |
|
||||
|
||||
### 2.8 Forward 独立限速写入链路补全(新增)
|
||||
|
||||
| 序号 | 任务 | 文件 | 状态 |
|
||||
|------|------|------|------|
|
||||
| N1 | forwardCreate 支持接收并校验可选 speedId,写入 Forward.SpeedID | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
| N2 | forwardUpdate 支持更新/清空 speedId,并触发服务重下发 | `go-backend/internal/http/handler/mutations.go` | ✅ 完成 |
|
||||
| N3 | CreateForwardTx 支持落库 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| N4 | UpdateForward 支持更新 speed_id | `go-backend/internal/store/repo/repository_mutations.go` | ✅ 完成 |
|
||||
| N5 | Forward 页面新增限速选择并透传 speedId | `vite-frontend/src/pages/forward.tsx` | ✅ 完成 |
|
||||
| N6 | Forward 相关契约测试补充 speedId 写入/清空断言 | `go-backend/tests/contract/forward_contract_test.go` | ✅ 完成 |
|
||||
|
||||
---
|
||||
|
||||
## 三、优先级说明
|
||||
|
||||
限速规则应用优先级:
|
||||
1. **Forward.SpeedID** - 转发级别的限速 (最高优先)
|
||||
2. **UserTunnel.SpeedID** - 用户隧道权限级别的限速 (回退)
|
||||
|
||||
---
|
||||
|
||||
## 四、数据库兼容性
|
||||
|
||||
- SpeedLimit 表: `tunnel_id` 和 `tunnel_name` 字段改为可空 (GORM AutoMigrate 自动处理)
|
||||
- Forward 表: 新增 `speed_id` 可空字段 (GORM AutoMigrate 自动处理)
|
||||
|
||||
---
|
||||
|
||||
## 五、验证检查项
|
||||
|
||||
### 5.1 功能验证(审计后)
|
||||
|
||||
- [x] 创建不限速规则的限速 (不绑定隧道)
|
||||
- [x] 创建绑定隧道的限速 (兼容旧逻辑)
|
||||
- [x] 编辑限速规则,切换隧道绑定状态
|
||||
- [ ] 删除限速规则
|
||||
- [ ] 转发列表正确显示 speedId
|
||||
|
||||
### 5.2 API 验证(审计后)
|
||||
|
||||
- [x] GET /api/speed-limit/list 返回可选 tunnelId
|
||||
- [x] POST /api/speed-limit/create 接受可选 tunnelId
|
||||
- [x] POST /api/speed-limit/update 接受可选 tunnelId
|
||||
- [ ] GET /api/forward/list 返回 speedId
|
||||
|
||||
### 5.3 兼容性验证(审计后)
|
||||
|
||||
- [x] 现有绑定隧道的限速规则继续正常工作
|
||||
- [ ] 现有 UserTunnel 的限速继续正常工作
|
||||
- [ ] 备份/恢复功能正常
|
||||
|
||||
### 5.4 Forward 独立限速闭环验证(新增)
|
||||
|
||||
- [x] POST /api/forward/create 接受 speedId 并写入 `forward.speed_id`
|
||||
- [x] POST /api/forward/update 可更新/清空 speedId
|
||||
- [x] Forward 表单可选择限速并提交 speedId
|
||||
- [ ] `syncForwardServices` 实际使用 Forward.SpeedID 而非仅回退 UserTunnel.SpeedID
|
||||
+12
-8
@@ -2,7 +2,7 @@
|
||||
|
||||
## 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 +17,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 +37,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 +49,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 +61,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
|
||||
```
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
# BACKEND HTTP HANDLER KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Generated:** Thu Feb 26 2026
|
||||
|
||||
## 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 +14,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 +26,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.).
|
||||
|
||||
@@ -152,11 +152,32 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
||||
return errors.New("转发入口端口不存在")
|
||||
}
|
||||
|
||||
userTunnelID, limiterID, speed, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
// Determine limiter from forward's SpeedID first, fallback to UserTunnel's limiter
|
||||
var limiterID *int64
|
||||
var speed *int
|
||||
|
||||
if forward.SpeedID.Valid && forward.SpeedID.Int64 > 0 {
|
||||
// Forward has its own speed limit
|
||||
speedVal, err := h.repo.GetSpeedLimitSpeed(forward.SpeedID.Int64)
|
||||
if err == nil && speedVal > 0 {
|
||||
limiterID = &forward.SpeedID.Int64
|
||||
speed = &speedVal
|
||||
}
|
||||
}
|
||||
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
|
||||
|
||||
if limiterID == nil {
|
||||
// Fall back to UserTunnel speed limit
|
||||
var utLimiterID *int64
|
||||
var utSpeed *int
|
||||
_, utLimiterID, utSpeed, err = h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
limiterID = utLimiterID
|
||||
speed = utSpeed
|
||||
}
|
||||
|
||||
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, 0)
|
||||
tunnelTLSProtocol, err := h.isTunnelSelectedTLSProtocol(forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -164,7 +185,9 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
||||
|
||||
for _, fp := range ports {
|
||||
if limiterID != nil && speed != nil {
|
||||
h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed)
|
||||
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(fp.NodeID)
|
||||
@@ -1030,12 +1053,16 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) {
|
||||
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {
|
||||
rate := float64(speed) / 8.0
|
||||
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
||||
payload := map[string]interface{}{
|
||||
"name": strconv.FormatInt(limiterID, 10),
|
||||
"limits": []string{limitStr},
|
||||
}
|
||||
_, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false)
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
|
||||
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1048,21 +1048,55 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
|
||||
if utErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, utErr.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
_, oldFlow, oldNum, oldExpTime, oldFlowReset, oldSpeedID, oldStatus, oldErr :=
|
||||
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
if oldErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, oldErr.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserTunnel(id,
|
||||
asInt64(req["flow"], 0),
|
||||
asInt(req["num"], 0),
|
||||
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
|
||||
asInt64(req["flowResetTime"], 1),
|
||||
nullableInt(asAnyToInt64Ptr(req["speedId"])),
|
||||
nullableInt(speedID),
|
||||
asInt(req["status"], 1),
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
|
||||
if utErr == nil {
|
||||
h.syncUserTunnelForwards(userID, tunnelID)
|
||||
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||
rollbackErr := h.repo.UpdateUserTunnel(
|
||||
id,
|
||||
oldFlow,
|
||||
int(oldNum),
|
||||
oldExpTime,
|
||||
oldFlowReset,
|
||||
oldSpeedID,
|
||||
oldStatus,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败,已回滚: %v", syncErr)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
@@ -1103,6 +1137,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
|
||||
return
|
||||
}
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if speedID != nil {
|
||||
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
|
||||
if speedErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
|
||||
return
|
||||
}
|
||||
if !exists {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
port = h.pickTunnelPort(tunnelID)
|
||||
@@ -1127,7 +1173,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if userName == "" {
|
||||
userName = "user"
|
||||
}
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port)
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, nullableInt(speedID))
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1202,6 +1248,24 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if strategy == "" {
|
||||
strategy = forward.Strategy
|
||||
}
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if speedID != nil {
|
||||
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
|
||||
if speedErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
|
||||
return
|
||||
}
|
||||
if !exists {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
newSpeedID := forward.SpeedID
|
||||
if speedID != nil {
|
||||
newSpeedID = sql.NullInt64{Int64: *speedID, Valid: true}
|
||||
} else if _, ok := req["speedId"]; ok {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
@@ -1225,7 +1289,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now); err != nil {
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -1586,29 +1650,37 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
tunnelID := asInt64(req["tunnelId"], 0)
|
||||
if tunnelID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
name := asString(req["name"])
|
||||
if name == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("名称不能为空"))
|
||||
return
|
||||
}
|
||||
tunnelName := h.repo.GetTunnelNameByID(tunnelID)
|
||||
if tunnelName == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||
return
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
speed := asInt(req["speed"], 100)
|
||||
|
||||
var tunnelID *int64
|
||||
var tunnelName string
|
||||
if tid := asInt64(req["tunnelId"], 0); tid > 0 {
|
||||
tunnelID = &tid
|
||||
tunnelName = h.repo.GetTunnelNameByID(tid)
|
||||
if tunnelName == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
id, err := h.repo.CreateSpeedLimit(name, speed, tunnelID, tunnelName, now, asInt(req["status"], 1))
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
_ = h.sendLimiterConfig(id, speed, tunnelID)
|
||||
|
||||
if tunnelID != nil && *tunnelID > 0 {
|
||||
_ = h.sendLimiterConfig(id, speed, *tunnelID)
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -1618,23 +1690,41 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
|
||||
id := asInt64(req["id"], 0)
|
||||
tunnelID := asInt64(req["tunnelId"], 0)
|
||||
if id <= 0 || tunnelID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
if id <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则ID不能为空"))
|
||||
return
|
||||
}
|
||||
tunnelName := h.repo.GetTunnelNameByID(tunnelID)
|
||||
if tunnelName == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||
|
||||
name := asString(req["name"])
|
||||
if name == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("名称不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
speed := asInt(req["speed"], 100)
|
||||
if err := h.repo.UpdateSpeedLimit(id, asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil {
|
||||
|
||||
var tunnelID *int64
|
||||
var tunnelName string
|
||||
if tid := asInt64(req["tunnelId"], 0); tid > 0 {
|
||||
tunnelID = &tid
|
||||
tunnelName = h.repo.GetTunnelNameByID(tid)
|
||||
if tunnelName == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateSpeedLimit(id, name, speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
_ = h.sendLimiterConfig(id, speed, tunnelID)
|
||||
|
||||
if tunnelID != nil && *tunnelID > 0 {
|
||||
_ = h.sendLimiterConfig(id, speed, *tunnelID)
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -1643,15 +1733,18 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
tunnelID := h.repo.GetSpeedLimitTunnelID(id)
|
||||
|
||||
if err := h.repo.DeleteSpeedLimit(id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if tunnelID > 0 {
|
||||
_ = h.sendDeleteLimiterConfig(id, tunnelID)
|
||||
|
||||
if tunnelID.Valid && tunnelID.Int64 > 0 {
|
||||
_ = h.sendDeleteLimiterConfig(id, tunnelID.Int64)
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2713,7 +2806,7 @@ func pickNodeAddressV6(node *nodeRecord) string {
|
||||
func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[string]interface{}) error {
|
||||
allocated := map[int64]int{}
|
||||
inNodes := asMapSlice(req["inNodeId"])
|
||||
for _, n := range inNodes {
|
||||
for i, n := range inNodes {
|
||||
nodeID := asInt64(n["nodeId"], 0)
|
||||
if nodeID <= 0 {
|
||||
continue
|
||||
@@ -2725,13 +2818,13 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str
|
||||
nodeID,
|
||||
sql.NullInt64{},
|
||||
defaultString(asString(n["strategy"]), "round"),
|
||||
0,
|
||||
i+1,
|
||||
defaultString(asString(n["protocol"]), "tls"),
|
||||
); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, n := range asMapSlice(req["outNodeId"]) {
|
||||
for i, n := range asMapSlice(req["outNodeId"]) {
|
||||
nodeID := asInt64(n["nodeId"], 0)
|
||||
if nodeID <= 0 {
|
||||
continue
|
||||
@@ -2751,7 +2844,7 @@ func (h *Handler) replaceTunnelChainsTx(tx *gorm.DB, tunnelID int64, req map[str
|
||||
nodeID,
|
||||
sql.NullInt64{Int64: int64(port), Valid: true},
|
||||
defaultString(asString(n["strategy"]), "round"),
|
||||
0,
|
||||
i+1,
|
||||
defaultString(asString(n["protocol"]), "tls"),
|
||||
); err != nil {
|
||||
return err
|
||||
@@ -2977,6 +3070,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
||||
h.repo.RollbackForwardFields(
|
||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||
oldForward.SpeedID,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
@@ -2998,6 +3092,10 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
reqFlow := asInt64(req["flow"], -1)
|
||||
reqNum := asInt(req["num"], -1)
|
||||
reqExpTime := asInt64(req["expTime"], -1)
|
||||
@@ -3038,7 +3136,24 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
reqStatus = 1
|
||||
}
|
||||
|
||||
return h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus)
|
||||
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||
insertedID, _, _, _, _, _, _, lookupErr := h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
if lookupErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚查询错误: %w", syncErr, lookupErr)
|
||||
}
|
||||
|
||||
if rollbackErr := h.repo.DeleteUserTunnel(insertedID); rollbackErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚删除错误: %w", syncErr, rollbackErr)
|
||||
}
|
||||
|
||||
return fmt.Errorf("下发失败,已回滚: %w", syncErr)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -3076,25 +3191,61 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
err = h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus)
|
||||
|
||||
if err == nil {
|
||||
h.syncUserTunnelForwards(userID, tunnelID)
|
||||
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
|
||||
return err
|
||||
}
|
||||
return err
|
||||
|
||||
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
|
||||
rollbackErr := h.repo.UpdateUserTunnelFields(
|
||||
existingID,
|
||||
currentSpeedID,
|
||||
currentFlow,
|
||||
int(currentNum),
|
||||
currentExpTime,
|
||||
currentFlowReset,
|
||||
currentStatus,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr)
|
||||
}
|
||||
|
||||
return fmt.Errorf("下发失败,已回滚: %w", syncErr)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) {
|
||||
func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error {
|
||||
forwards, err := h.listForwardsByTunnel(tunnelID)
|
||||
if err != nil {
|
||||
return
|
||||
return err
|
||||
}
|
||||
for i := range forwards {
|
||||
f := &forwards[i]
|
||||
if f.UserID == userID {
|
||||
_ = h.syncForwardServices(f, "UpdateService", true)
|
||||
if err := h.syncForwardServices(f, "UpdateService", true); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateSpeedLimitReference(speedID *int64) error {
|
||||
if speedID == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
exists, err := h.repo.SpeedLimitExists(*speedID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !exists {
|
||||
return errors.New("限速规则不存在")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func asAnySlice(v interface{}) []interface{} {
|
||||
|
||||
@@ -29,19 +29,20 @@ 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" }
|
||||
@@ -83,14 +84,14 @@ type Node struct {
|
||||
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" }
|
||||
@@ -395,6 +396,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 +423,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 +494,7 @@ type ForwardRecord struct {
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
}
|
||||
|
||||
// TunnelRecord is a minimal tunnel view used by control plane.
|
||||
|
||||
@@ -656,7 +656,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var users []model.User
|
||||
if err := r.db.Where("role_id != ?", 0).Order("id ASC").Find(&users).Error; err != nil {
|
||||
if err := r.db.Where("role_id != ?", 0).Order("id DESC").Find(&users).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(users))
|
||||
@@ -678,17 +678,23 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var limits []model.SpeedLimit
|
||||
if err := r.db.Order("id ASC").Find(&limits).Error; err != nil {
|
||||
if err := r.db.Order("id DESC").Find(&limits).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]map[string]interface{}, 0, len(limits))
|
||||
for _, sl := range limits {
|
||||
items = append(items, map[string]interface{}{
|
||||
item := map[string]interface{}{
|
||||
"id": sl.ID, "name": sl.Name, "speed": sl.Speed,
|
||||
"tunnelId": sl.TunnelID, "tunnelName": sl.TunnelName,
|
||||
"status": sl.Status, "createdTime": sl.CreatedTime,
|
||||
"updatedTime": nullableInt64(sl.UpdatedTime),
|
||||
})
|
||||
}
|
||||
if sl.TunnelID.Valid {
|
||||
item["tunnelId"] = sl.TunnelID.Int64
|
||||
}
|
||||
if sl.TunnelName.Valid {
|
||||
item["tunnelName"] = sl.TunnelName.String
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
@@ -712,11 +718,12 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
CreatedTime int64
|
||||
Status int
|
||||
Inx int
|
||||
SpeedID sql.NullInt64
|
||||
}
|
||||
|
||||
var rows []fwdRow
|
||||
err := r.db.Model(&model.Forward{}).
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx").
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id").
|
||||
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
|
||||
Order("forward.inx ASC, forward.id ASC").
|
||||
Find(&rows).Error
|
||||
@@ -730,14 +737,18 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, map[string]interface{}{
|
||||
item := map[string]interface{}{
|
||||
"id": row.ID, "userId": row.UserID, "userName": row.UserName,
|
||||
"name": row.Name, "tunnelId": row.TunnelID, "tunnelName": row.TunnelName,
|
||||
"inIp": nullableForwardIngress(inIP), "inPort": nullableInt64(inPort),
|
||||
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
|
||||
"inFlow": row.InFlow, "outFlow": row.OutFlow,
|
||||
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
|
||||
})
|
||||
}
|
||||
if row.SpeedID.Valid {
|
||||
item["speedId"] = row.SpeedID.Int64
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
@@ -1308,7 +1319,6 @@ func (r *Repository) ListActiveForwardPeerShareRuntimesByNodeAndServiceName(node
|
||||
return items, nil
|
||||
}
|
||||
|
||||
|
||||
func (r *Repository) ListActiveForwardPeerShareRuntimeServiceNamesByNode(nodeID int64) ([]string, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
@@ -1813,9 +1823,15 @@ func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) {
|
||||
for _, sl := range sls {
|
||||
b := model.SpeedLimitBackup{
|
||||
ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed),
|
||||
TunnelID: sl.TunnelID, TunnelName: sl.TunnelName,
|
||||
CreatedTime: sl.CreatedTime, Status: sl.Status,
|
||||
}
|
||||
if sl.TunnelID.Valid {
|
||||
tid := sl.TunnelID.Int64
|
||||
b.TunnelID = &tid
|
||||
}
|
||||
if sl.TunnelName.Valid {
|
||||
b.TunnelName = sl.TunnelName.String
|
||||
}
|
||||
if sl.UpdatedTime.Valid {
|
||||
b.UpdatedTime = sl.UpdatedTime.Int64
|
||||
}
|
||||
@@ -2186,12 +2202,18 @@ func importSpeedLimits(tx *gorm.DB, speedLimits []model.SpeedLimitBackup, now in
|
||||
ID: sl.ID,
|
||||
Name: sl.Name,
|
||||
Speed: int(sl.Speed),
|
||||
TunnelID: sl.TunnelID,
|
||||
TunnelName: sl.TunnelName,
|
||||
TunnelID: sql.NullInt64{Int64: 0, Valid: false},
|
||||
TunnelName: sql.NullString{String: "", Valid: false},
|
||||
CreatedTime: sl.CreatedTime,
|
||||
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
|
||||
Status: sl.Status,
|
||||
}
|
||||
if sl.TunnelID != nil {
|
||||
item.TunnelID = sql.NullInt64{Int64: *sl.TunnelID, Valid: true}
|
||||
}
|
||||
if sl.TunnelName != "" {
|
||||
item.TunnelName = sql.NullString{String: sl.TunnelName, Valid: true}
|
||||
}
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
|
||||
@@ -46,6 +46,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 {
|
||||
|
||||
@@ -228,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
|
||||
@@ -251,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) {
|
||||
@@ -260,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,7 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -99,6 +101,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"
|
||||
@@ -169,3 +172,15 @@ func (r *Repository) SpeedLimitExists(id int64) (bool, error) {
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
@@ -657,7 +657,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")
|
||||
}
|
||||
@@ -668,6 +668,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
|
||||
}
|
||||
@@ -724,7 +725,7 @@ func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct {
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, now int64) {
|
||||
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
|
||||
}
|
||||
@@ -738,6 +739,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
|
||||
}
|
||||
@@ -764,51 +766,66 @@ 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, tunnelID *int64, tunnelName string, 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,
|
||||
}
|
||||
if tunnelID != nil {
|
||||
sl.TunnelID = sql.NullInt64{Int64: *tunnelID, Valid: true}
|
||||
}
|
||||
if tunnelName != "" {
|
||||
sl.TunnelName = sql.NullString{String: tunnelName, Valid: true}
|
||||
}
|
||||
if err := r.db.Create(&sl).Error; err != nil {
|
||||
return 0, err
|
||||
}
|
||||
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, tunnelID *int64, tunnelName string, 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,
|
||||
"updated_time": sql.NullInt64{
|
||||
Int64: now,
|
||||
Valid: true,
|
||||
},
|
||||
}
|
||||
if tunnelID != nil {
|
||||
updates["tunnel_id"] = sql.NullInt64{Int64: *tunnelID, Valid: true}
|
||||
} else {
|
||||
updates["tunnel_id"] = sql.NullInt64{Int64: 0, Valid: false}
|
||||
}
|
||||
if tunnelName != "" {
|
||||
updates["tunnel_name"] = sql.NullString{String: tunnelName, Valid: true}
|
||||
} else {
|
||||
updates["tunnel_name"] = sql.NullString{String: "", Valid: false}
|
||||
}
|
||||
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
|
||||
Updates(updates).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) int64 {
|
||||
func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) sql.NullInt64 {
|
||||
if r == nil || r.db == nil {
|
||||
return 0
|
||||
return sql.NullInt64{Valid: false}
|
||||
}
|
||||
var sl model.SpeedLimit
|
||||
if err := r.db.Select("tunnel_id").Where("id = ?", speedLimitID).First(&sl).Error; err != nil {
|
||||
return 0
|
||||
return sql.NullInt64{Valid: false}
|
||||
}
|
||||
return sl.TunnelID
|
||||
}
|
||||
@@ -1190,7 +1207,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, speedID interface{}) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
@@ -1209,6 +1226,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
|
||||
|
||||
@@ -2,6 +2,7 @@ package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -471,6 +472,144 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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 jsonNumber(v int64) string {
|
||||
return strconv.FormatInt(v, 10)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,402 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
)
|
||||
|
||||
func TestForwardCreateRollbackWhenLimiterDispatchFailsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := 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 := r.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "limiter-fail-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "limiter-fail-tunnel")
|
||||
|
||||
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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "limiter-fail-node", "limiter-fail-secret", "10.20.0.1", "10.20.0.1", "", "32000-32010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "limiter-fail-node")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 32001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "limiter-fail-rule", 1024, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, r, "limiter-fail-rule")
|
||||
|
||||
stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "limiter-fail-secret", map[string]string{
|
||||
"addlimiters": "mock add limiters failed",
|
||||
})
|
||||
defer stopNode()
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"name": "limiter-fail-forward",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedID,
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", 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 create failure on limiter dispatch, got code=0")
|
||||
}
|
||||
|
||||
forwardCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM forward WHERE name = ?`, "limiter-fail-forward")
|
||||
if forwardCount != 0 {
|
||||
t.Fatalf("expected forward rollback delete on limiter failure, got count=%d", forwardCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBatchAssignRollbackWhenLimiterDispatchFailsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := 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 := 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, 'assign_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)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "assign-limiter-fail-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "assign-limiter-fail-tunnel")
|
||||
|
||||
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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "assign-limiter-fail-node", "assign-limiter-fail-secret", "10.21.0.1", "10.21.0.1", "", "33000-33010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "assign-limiter-fail-node")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 33001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "assign-limiter-fail-rule", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, r, "assign-limiter-fail-rule")
|
||||
|
||||
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(21, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %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(2, 'assign_user', 'assign-limiter-fail-forward', ?, '9.9.9.9:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, "assign-limiter-fail-forward")
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "assign-limiter-fail-secret", map[string]string{
|
||||
"addlimiters": "mock add limiters failed",
|
||||
})
|
||||
defer stopNode()
|
||||
|
||||
assignPayload := map[string]interface{}{
|
||||
"userId": 2,
|
||||
"tunnels": []map[string]interface{}{{
|
||||
"tunnelId": tunnelID,
|
||||
"speedId": speedID,
|
||||
}},
|
||||
}
|
||||
body, err := json.Marshal(assignPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal assign payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", 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 assign failure on limiter dispatch, got code=0")
|
||||
}
|
||||
|
||||
var persistedSpeedID sql.NullInt64
|
||||
if err := r.DB().Raw(`SELECT speed_id FROM user_tunnel WHERE user_id = 2 AND tunnel_id = ?`, tunnelID).Row().Scan(&persistedSpeedID); err != nil {
|
||||
t.Fatalf("query user_tunnel speed_id: %v", err)
|
||||
}
|
||||
if persistedSpeedID.Valid {
|
||||
t.Fatalf("expected speed_id rollback to NULL, got %d", persistedSpeedID.Int64)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBatchAssignInsertRollbackWhenLimiterDispatchFailsContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, r := 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 := 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, 'assign_insert_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)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "assign-insert-limiter-fail-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, r, "assign-insert-limiter-fail-tunnel")
|
||||
|
||||
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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "assign-insert-limiter-fail-node", "assign-insert-limiter-fail-secret", "10.22.0.1", "10.22.0.1", "", "34000-34010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, r, "assign-insert-limiter-fail-node")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 34001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "assign-insert-limiter-fail-rule", 3072, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, r, "assign-insert-limiter-fail-rule")
|
||||
|
||||
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(3, 'assign_insert_user', 'assign-insert-limiter-fail-forward', ?, '8.8.4.4:53', 'fifo', 0, 0, ?, ?, 1, 0)
|
||||
`, tunnelID, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, "assign-insert-limiter-fail-forward")
|
||||
|
||||
if err := r.DB().Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, 34001).Error; err != nil {
|
||||
t.Fatalf("insert forward_port: %v", err)
|
||||
}
|
||||
|
||||
stopNode := startMockNodeSessionWithCommandFailures(t, server.URL, "assign-insert-limiter-fail-secret", map[string]string{
|
||||
"addlimiters": "mock add limiters failed",
|
||||
})
|
||||
defer stopNode()
|
||||
|
||||
assignPayload := map[string]interface{}{
|
||||
"userId": 3,
|
||||
"tunnels": []map[string]interface{}{{
|
||||
"tunnelId": tunnelID,
|
||||
"speedId": speedID,
|
||||
}},
|
||||
}
|
||||
body, err := json.Marshal(assignPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal assign payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", 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 assign(insert) failure on limiter dispatch, got code=0")
|
||||
}
|
||||
|
||||
insertedCount := mustQueryInt(t, r, `SELECT COUNT(1) FROM user_tunnel WHERE user_id = 3 AND tunnel_id = ?`, tunnelID)
|
||||
if insertedCount != 0 {
|
||||
t.Fatalf("expected inserted user_tunnel rollback delete, got count=%d", insertedCount)
|
||||
}
|
||||
}
|
||||
|
||||
func startMockNodeSessionWithCommandFailures(t *testing.T, baseURL string, nodeSecret string, failCommands map[string]string) func() {
|
||||
t.Helper()
|
||||
|
||||
u, err := url.Parse(baseURL)
|
||||
if err != nil {
|
||||
t.Fatalf("parse provider url: %v", err)
|
||||
}
|
||||
if strings.EqualFold(u.Scheme, "https") {
|
||||
u.Scheme = "wss"
|
||||
} else {
|
||||
u.Scheme = "ws"
|
||||
}
|
||||
u.Path = "/system-info"
|
||||
q := u.Query()
|
||||
q.Set("type", "1")
|
||||
q.Set("secret", nodeSecret)
|
||||
q.Set("version", "v1")
|
||||
q.Set("http", "1")
|
||||
q.Set("tls", "1")
|
||||
q.Set("socks", "1")
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
conn, _, err := websocket.DefaultDialer.Dial(u.String(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("dial mock node websocket: %v", err)
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for {
|
||||
_, raw, readErr := conn.ReadMessage()
|
||||
if readErr != nil {
|
||||
return
|
||||
}
|
||||
|
||||
plain := raw
|
||||
var wrap struct {
|
||||
Encrypted bool `json:"encrypted"`
|
||||
Data string `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &wrap); err == nil && wrap.Encrypted && strings.TrimSpace(wrap.Data) != "" {
|
||||
crypto, cryptoErr := security.NewAESCrypto(nodeSecret)
|
||||
if cryptoErr == nil {
|
||||
if dec, decErr := crypto.Decrypt(wrap.Data); decErr == nil {
|
||||
plain = []byte(dec)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var cmd struct {
|
||||
Type string `json:"type"`
|
||||
RequestID string `json:"requestId"`
|
||||
}
|
||||
if err := json.Unmarshal(plain, &cmd); err != nil {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(cmd.RequestID) == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
cmdType := strings.TrimSpace(cmd.Type)
|
||||
failMsg, shouldFail := failCommands[strings.ToLower(cmdType)]
|
||||
|
||||
respType := fmt.Sprintf("%sResponse", cmdType)
|
||||
respPayload := map[string]interface{}{
|
||||
"type": respType,
|
||||
"success": !shouldFail,
|
||||
"message": "OK",
|
||||
"requestId": cmd.RequestID,
|
||||
}
|
||||
if shouldFail {
|
||||
if strings.TrimSpace(failMsg) == "" {
|
||||
failMsg = "mock command failed"
|
||||
}
|
||||
respPayload["message"] = failMsg
|
||||
}
|
||||
|
||||
respBytes, err := json.Marshal(respPayload)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
_ = conn.WriteMessage(websocket.TextMessage, respBytes)
|
||||
}
|
||||
}()
|
||||
|
||||
var stopOnce sync.Once
|
||||
return func() {
|
||||
stopOnce.Do(func() {
|
||||
_ = conn.Close()
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,462 @@
|
||||
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"
|
||||
)
|
||||
|
||||
// TestSpeedLimitWithoutTunnelContract tests that speed limits can be created without binding to a tunnel
|
||||
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)
|
||||
}
|
||||
|
||||
// Create a speed limit without tunnel binding
|
||||
t.Run("create speed limit without tunnel", 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)
|
||||
})
|
||||
|
||||
// Verify the speed limit has null tunnelId
|
||||
t.Run("list speed limits shows null tunnelId", 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)
|
||||
}
|
||||
|
||||
// Find our speed limit
|
||||
var found bool
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if m["name"] == "test-limit-no-tunnel" {
|
||||
found = true
|
||||
// tunnelId should be nil/not present for unbound speed limits
|
||||
if tunnelID, exists := m["tunnelId"]; exists && tunnelID != nil {
|
||||
t.Fatalf("expected tunnelId to be nil for unbound speed limit, got %v", tunnelID)
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Fatal("speed limit 'test-limit-no-tunnel' not found in list")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSpeedLimitWithTunnelContract tests that speed limits can still be bound to tunnels
|
||||
func TestSpeedLimitWithTunnelContract(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)
|
||||
}
|
||||
|
||||
// First create a tunnel
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-for-limit")
|
||||
|
||||
// Create a speed limit with tunnel binding
|
||||
t.Run("create speed limit with tunnel", func(t *testing.T) {
|
||||
body := `{"name":"test-limit-with-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)
|
||||
})
|
||||
|
||||
// Verify the speed limit has the tunnelId
|
||||
t.Run("list speed limits shows tunnelId", 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)
|
||||
}
|
||||
|
||||
var found bool
|
||||
for _, item := range data {
|
||||
m, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if m["name"] == "test-limit-with-tunnel" {
|
||||
found = true
|
||||
tunnelIDVal, exists := m["tunnelId"]
|
||||
if !exists || tunnelIDVal == nil {
|
||||
t.Fatal("expected tunnelId to be present for bound speed limit")
|
||||
}
|
||||
// Verify tunnelId matches
|
||||
if tunnelIDFloat, ok := tunnelIDVal.(float64); ok {
|
||||
if int64(tunnelIDFloat) != tunnelID {
|
||||
t.Fatalf("expected tunnelId %d, got %d", tunnelID, int64(tunnelIDFloat))
|
||||
}
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !found {
|
||||
t.Fatal("speed limit 'test-limit-with-tunnel' not found in list")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSpeedLimitUpdateTunnelBindingContract tests updating speed limit tunnel binding
|
||||
func TestSpeedLimitUpdateTunnelBindingContract(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)
|
||||
}
|
||||
|
||||
// Create a tunnel
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "test-tunnel-update")
|
||||
|
||||
// Create a speed limit without tunnel
|
||||
speedLimitID := mustCreateSpeedLimitRepo(t, r, "test-limit-update", 0)
|
||||
|
||||
// Update to bind to tunnel
|
||||
t.Run("update speed limit to bind tunnel", func(t *testing.T) {
|
||||
body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update","speed":150,"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)
|
||||
})
|
||||
|
||||
// Verify binding
|
||||
t.Run("verify tunnel binding after update", 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-update" {
|
||||
tunnelIDVal, exists := m["tunnelId"]
|
||||
if !exists || tunnelIDVal == nil {
|
||||
t.Fatal("expected tunnelId to be present after update")
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("speed limit 'test-limit-update' not found")
|
||||
})
|
||||
|
||||
// Update to unbind from tunnel (set tunnelId to null)
|
||||
t.Run("update speed limit to unbind tunnel", func(t *testing.T) {
|
||||
body := `{"id":` + jsonInt(speedLimitID) + `,"name":"test-limit-update","speed":150,"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)
|
||||
})
|
||||
|
||||
// Verify unbinding
|
||||
t.Run("verify tunnel unbinding after update", 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-update" {
|
||||
if tunnelIDVal, exists := m["tunnelId"]; exists && tunnelIDVal != nil {
|
||||
t.Fatalf("expected tunnelId to be nil after unbinding, got %v", tunnelIDVal)
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("speed limit 'test-limit-update' not found")
|
||||
})
|
||||
}
|
||||
|
||||
// TestSpeedLimitDatabaseNullableFields tests database-level nullable fields
|
||||
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() })
|
||||
|
||||
// Create speed limit via repository
|
||||
t.Run("repository create speed limit without tunnel", func(t *testing.T) {
|
||||
id, err := r.CreateSpeedLimit("db-test-limit", 100, nil, "", 1, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateSpeedLimit failed: %v", err)
|
||||
}
|
||||
if id <= 0 {
|
||||
t.Fatalf("expected valid id, got %d", id)
|
||||
}
|
||||
})
|
||||
|
||||
// Verify TunnelID is null in database
|
||||
t.Run("verify null TunnelID in database", func(t *testing.T) {
|
||||
var tunnelID sql.NullInt64
|
||||
var tunnelName sql.NullString
|
||||
err := r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?", "db-test-limit").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)
|
||||
}
|
||||
})
|
||||
|
||||
// Create a tunnel for binding test
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "db-test-tunnel")
|
||||
|
||||
// Create speed limit with tunnel
|
||||
t.Run("repository create speed limit with tunnel", func(t *testing.T) {
|
||||
id, err := r.CreateSpeedLimit("db-test-limit-with-tunnel", 200, &tunnelID, "db-test-tunnel", 1, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateSpeedLimit failed: %v", err)
|
||||
}
|
||||
if id <= 0 {
|
||||
t.Fatalf("expected valid id, got %d", id)
|
||||
}
|
||||
})
|
||||
|
||||
// Verify TunnelID is set
|
||||
t.Run("verify TunnelID is set in database", func(t *testing.T) {
|
||||
var dbTunnelID sql.NullInt64
|
||||
var dbTunnelName sql.NullString
|
||||
err := r.DB().Raw("SELECT tunnel_id, tunnel_name FROM speed_limit WHERE name = ?", "db-test-limit-with-tunnel").Row().Scan(&dbTunnelID, &dbTunnelName)
|
||||
if err != nil {
|
||||
t.Fatalf("query failed: %v", err)
|
||||
}
|
||||
if !dbTunnelID.Valid {
|
||||
t.Fatal("expected TunnelID to be valid")
|
||||
}
|
||||
if dbTunnelID.Int64 != tunnelID {
|
||||
t.Fatalf("expected TunnelID %d, got %d", tunnelID, dbTunnelID.Int64)
|
||||
}
|
||||
if !dbTunnelName.Valid || dbTunnelName.String != "db-test-tunnel" {
|
||||
t.Fatalf("expected TunnelName 'db-test-tunnel', got %v", dbTunnelName.String)
|
||||
}
|
||||
})
|
||||
|
||||
// Test GetSpeedLimitTunnelID returns correct nullability
|
||||
t.Run("GetSpeedLimitTunnelID returns null for unbound limit", func(t *testing.T) {
|
||||
result := r.GetSpeedLimitTunnelID(1) // First speed limit (db-test-limit)
|
||||
if result.Valid {
|
||||
t.Fatalf("expected GetSpeedLimitTunnelID to return invalid/null, got valid with value %d", result.Int64)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("GetSpeedLimitTunnelID returns value for bound limit", func(t *testing.T) {
|
||||
result := r.GetSpeedLimitTunnelID(2) // Second speed limit (db-test-limit-with-tunnel)
|
||||
if !result.Valid {
|
||||
t.Fatal("expected GetSpeedLimitTunnelID to return valid result for bound limit")
|
||||
}
|
||||
if result.Int64 != tunnelID {
|
||||
t.Fatalf("expected TunnelID %d, got %d", tunnelID, result.Int64)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSpeedLimitUpdateUnbindFromTunnel tests unbinding a speed limit from a tunnel
|
||||
func TestSpeedLimitUpdateUnbindFromTunnel(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "speed-limit-unbind.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
// Create tunnel
|
||||
tunnelID := mustCreateSpeedLimitTunnel(t, r, "unbind-test-tunnel")
|
||||
|
||||
// Create speed limit bound to tunnel
|
||||
speedLimitID, err := r.CreateSpeedLimit("unbind-test-limit", 300, &tunnelID, "unbind-test-tunnel", 1, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("create speed limit: %v", err)
|
||||
}
|
||||
|
||||
// Verify initial binding
|
||||
t.Run("verify initial binding", func(t *testing.T) {
|
||||
result := r.GetSpeedLimitTunnelID(speedLimitID)
|
||||
if !result.Valid {
|
||||
t.Fatal("expected initial binding to tunnel")
|
||||
}
|
||||
if result.Int64 != tunnelID {
|
||||
t.Fatalf("expected tunnel ID %d, got %d", tunnelID, result.Int64)
|
||||
}
|
||||
})
|
||||
|
||||
// Update to unbind
|
||||
t.Run("unbind speed limit from tunnel via UpdateSpeedLimit", func(t *testing.T) {
|
||||
err := r.UpdateSpeedLimit(speedLimitID, "unbind-test-limit", 300, nil, "", 1, time.Now().UnixMilli())
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateSpeedLimit failed: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
// Verify unbinding
|
||||
t.Run("verify unbinding after update", func(t *testing.T) {
|
||||
result := r.GetSpeedLimitTunnelID(speedLimitID)
|
||||
if result.Valid {
|
||||
t.Fatalf("expected GetSpeedLimitTunnelID to return invalid/null after unbind, got valid with value %d", result.Int64)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestSpeedLimitGetSpeed tests the GetSpeedLimitSpeed function
|
||||
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() })
|
||||
|
||||
// Create speed limit
|
||||
speedLimitID, err := r.CreateSpeedLimit("get-speed-test", 500, nil, "", 1, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("create speed limit: %v", err)
|
||||
}
|
||||
|
||||
// Test GetSpeedLimitSpeed
|
||||
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")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Helper functions
|
||||
|
||||
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, tunnelID int64) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
var tid *int64
|
||||
if tunnelID > 0 {
|
||||
tid = &tunnelID
|
||||
}
|
||||
id, err := r.CreateSpeedLimit(name, 100, tid, "", now, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("create speed limit failed: %v", err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
+7
-5
@@ -1,6 +1,6 @@
|
||||
# GO-GOST SERVICE KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Sun Feb 15 2026
|
||||
**Generated:** Thu Feb 26 2026
|
||||
|
||||
## OVERVIEW
|
||||
Forwarding agent built on GOST v3 with a local fork of `github.com/go-gost/x` under `x/`.
|
||||
@@ -19,16 +19,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/`.
|
||||
|
||||
+11
-10
@@ -6,27 +6,28 @@ Local fork of `github.com/go-gost/x` used by `go-gost/` via `replace github.com/
|
||||
## 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 +42,4 @@ go-gost/x/
|
||||
```bash
|
||||
cd go-gost/x
|
||||
go test ./...
|
||||
```
|
||||
```
|
||||
+11
-12
@@ -1,9 +1,9 @@
|
||||
# VITE FRONTEND KNOWLEDGE BASE
|
||||
|
||||
**Generated:** Thu Feb 19 2026
|
||||
**Commit:** 137c34e
|
||||
**Generated:** Thu Feb 26 2026
|
||||
**Commit:** 21008cc
|
||||
**Branch:** main
|
||||
**Tag:** 2.1.4-rc2
|
||||
**Tag:** 2.1.5-rc15
|
||||
|
||||
## OVERVIEW
|
||||
Web management console for FLVX.
|
||||
@@ -15,8 +15,8 @@ vite-frontend/
|
||||
├── src/
|
||||
│ ├── api/ # Axios wrapper + typed endpoint helpers
|
||||
│ ├── components/ui/ # shadcn/radix primitive components
|
||||
│ ├── shadcn-bridge/heroui/ # HeroUI-compatible facade used by pages/layouts
|
||||
│ ├── pages/ # Route views + page modules (forward/node/tunnel split helpers)
|
||||
│ ├── shadcn-bridge/heroui/ # HeroUI-compatible facade (23 components)
|
||||
│ ├── pages/ # Route views + page modules (forward/node/tunnel)
|
||||
│ ├── hooks/ # H5/WebView/mobile hooks
|
||||
│ ├── styles/
|
||||
│ │ ├── globals.css # Base styles + imports tailwind-theme.pcss
|
||||
@@ -25,8 +25,8 @@ vite-frontend/
|
||||
│ ├── main.tsx # ReactDOM + BrowserRouter + Provider
|
||||
│ └── provider.tsx # Toast/theme/provider composition
|
||||
├── components.json # shadcn/ui config
|
||||
├── tailwind.config.js # Compatibility config still used by migration scaffolding
|
||||
├── vite.config.ts # base '/', host 0.0.0.0:3000; build minify/treeshake disabled
|
||||
├── tailwind.config.js # Compatibility config for migration scaffolding
|
||||
├── vite.config.ts # base '/', host 0.0.0.0:3000; minify/treeshake disabled
|
||||
└── package.json
|
||||
```
|
||||
|
||||
@@ -47,8 +47,8 @@ vite-frontend/
|
||||
- **API Envelope**: Responses follow `{code, msg, data, ts}`.
|
||||
- **UI Imports**: Use `src/shadcn-bridge/heroui/*` in app pages/layouts for compatibility.
|
||||
- **Semantic Colors**: Keep `globals.css -> tailwind-theme.pcss` import intact or semantic classes break.
|
||||
- **Build profile**: `minify: false`, `treeshake: false` for easier debugging.
|
||||
- **Layout mode**: H5/mobile mode still controlled by existing route/query and hook logic.
|
||||
- **Build profile**: `minify: false`, `treeshake: false` for debugging.
|
||||
- **Layout mode**: H5/mobile mode controlled by existing route/query and hook logic.
|
||||
|
||||
## ANTI-PATTERNS
|
||||
- **DO NOT ADD** `Bearer` to auth header in frontend requests.
|
||||
@@ -57,10 +57,9 @@ vite-frontend/
|
||||
- **DO NOT ADD** frontend tests; no Vitest/Jest setup exists.
|
||||
|
||||
## NOTES
|
||||
- PR `#144` (shadcn migration) and PR `#142` (user-group binding) are merged in `main`.
|
||||
- Release tag `2.1.4-rc2` points to commit `137c34e`.
|
||||
- Button border/color parity depends on both bridge mapping and semantic Tailwind token export.
|
||||
- Uses `rolldown-vite` (experimental Rust bundler) instead of standard Vite.
|
||||
- Build outputs are non-minified (debugging mode).
|
||||
- No test infrastructure exists (Vitest/Jest not configured).
|
||||
|
||||
## COMMANDS
|
||||
```bash
|
||||
|
||||
@@ -51,6 +51,7 @@ export interface ForwardApiItem {
|
||||
outFlow?: number;
|
||||
userId?: number;
|
||||
tunnelId?: number;
|
||||
speedId?: number | null;
|
||||
inx?: number;
|
||||
[key: string]: unknown;
|
||||
}
|
||||
@@ -96,10 +97,10 @@ export interface StatisticsFlowApiItem {
|
||||
export interface SpeedLimitApiItem {
|
||||
id: number;
|
||||
name: string;
|
||||
tunnelId: number;
|
||||
tunnelId?: number | null;
|
||||
speed: number;
|
||||
status: number;
|
||||
tunnelName: string;
|
||||
tunnelName?: string;
|
||||
createdTime: string;
|
||||
updatedTime: string;
|
||||
uploadSpeed?: number;
|
||||
@@ -284,6 +285,7 @@ export interface ForwardMutationPayload {
|
||||
inPort?: number | null;
|
||||
remoteAddr?: string;
|
||||
strategy?: string;
|
||||
speedId?: number | null;
|
||||
}
|
||||
|
||||
export interface SpeedLimitMutationPayload {
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import type { SpeedLimitApiItem } from "@/api/types";
|
||||
|
||||
import { useState, useEffect, useMemo } from "react";
|
||||
import toast from "react-hot-toast";
|
||||
import {
|
||||
@@ -50,6 +52,7 @@ import { Checkbox } from "@/shadcn-bridge/heroui/checkbox";
|
||||
import {
|
||||
createForward,
|
||||
getForwardList,
|
||||
getSpeedLimitList,
|
||||
getPeerShareList,
|
||||
getPeerRemoteUsageList,
|
||||
updateForward,
|
||||
@@ -105,6 +108,7 @@ interface Forward {
|
||||
userName?: string;
|
||||
userId?: number;
|
||||
inx?: number;
|
||||
speedId?: number | null;
|
||||
}
|
||||
|
||||
interface Tunnel {
|
||||
@@ -123,12 +127,14 @@ interface ForwardForm {
|
||||
remoteAddr: string;
|
||||
interfaceName?: string;
|
||||
strategy: string;
|
||||
speedId: number | null;
|
||||
}
|
||||
|
||||
export default function ForwardPage() {
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [forwards, setForwards] = useState<Forward[]>([]);
|
||||
const [tunnels, setTunnels] = useState<Tunnel[]>([]);
|
||||
const [speedLimits, setSpeedLimits] = useState<SpeedLimitApiItem[]>([]);
|
||||
const isMobile = useMobileBreakpoint();
|
||||
const [searchKeyword, setSearchKeyword] = useLocalStorageState(
|
||||
"forward-search-keyword",
|
||||
@@ -206,6 +212,7 @@ export default function ForwardPage() {
|
||||
remoteAddr: "",
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
});
|
||||
|
||||
// 表单验证错误
|
||||
@@ -327,7 +334,9 @@ export default function ForwardPage() {
|
||||
|
||||
const resolveShareIdForForward = (forward: Forward): number | null => {
|
||||
const candidates = new Set<number>();
|
||||
const shareIdFromName = parseShareIdFromTunnelName(forward.tunnelName || "");
|
||||
const shareIdFromName = parseShareIdFromTunnelName(
|
||||
forward.tunnelName || "",
|
||||
);
|
||||
|
||||
if (shareIdFromName) {
|
||||
candidates.add(shareIdFromName);
|
||||
@@ -445,9 +454,10 @@ export default function ForwardPage() {
|
||||
const loadData = async (lod = true) => {
|
||||
setLoading(lod);
|
||||
try {
|
||||
const [forwardsRes, tunnelsRes] = await Promise.all([
|
||||
const [forwardsRes, tunnelsRes, speedLimitsRes] = await Promise.all([
|
||||
getForwardList(),
|
||||
userTunnel(),
|
||||
getSpeedLimitList(),
|
||||
]);
|
||||
|
||||
if (forwardsRes.code === 0) {
|
||||
@@ -481,6 +491,10 @@ export default function ForwardPage() {
|
||||
setTunnels(tunnelsRes.data || []);
|
||||
} else {
|
||||
}
|
||||
|
||||
if (speedLimitsRes.code === 0) {
|
||||
setSpeedLimits(speedLimitsRes.data || []);
|
||||
}
|
||||
} catch {
|
||||
toast.error("加载数据失败");
|
||||
} finally {
|
||||
@@ -489,6 +503,30 @@ export default function ForwardPage() {
|
||||
};
|
||||
|
||||
// 表单验证
|
||||
const noLimitSpeedLimitIds = useMemo(() => {
|
||||
return new Set(
|
||||
speedLimits
|
||||
.filter((speedLimit) => speedLimit.name.trim() === "不限速")
|
||||
.map((speedLimit) => speedLimit.id),
|
||||
);
|
||||
}, [speedLimits]);
|
||||
|
||||
const availableSpeedLimits = useMemo(() => {
|
||||
return speedLimits.filter(
|
||||
(speedLimit) => !noLimitSpeedLimitIds.has(speedLimit.id),
|
||||
);
|
||||
}, [speedLimits, noLimitSpeedLimitIds]);
|
||||
|
||||
const normalizeSpeedId = (speedId?: number | null): number | null => {
|
||||
if (speedId === null || speedId === undefined) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return noLimitSpeedLimitIds.has(speedId) ? null : speedId;
|
||||
};
|
||||
|
||||
const selectedSpeedId = normalizeSpeedId(form.speedId);
|
||||
|
||||
const validateForm = (): boolean => {
|
||||
const newErrors: { [key: string]: string } = {};
|
||||
|
||||
@@ -555,6 +593,7 @@ export default function ForwardPage() {
|
||||
remoteAddr: "",
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -572,6 +611,7 @@ export default function ForwardPage() {
|
||||
remoteAddr: forward.remoteAddr.split(",").join("\n"),
|
||||
interfaceName: forward.interfaceName || "",
|
||||
strategy: forward.strategy || "fifo",
|
||||
speedId: normalizeSpeedId(forward.speedId),
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -651,6 +691,7 @@ export default function ForwardPage() {
|
||||
inPort: form.inPort,
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizeSpeedId(form.speedId),
|
||||
};
|
||||
|
||||
res = await updateForward(updateData);
|
||||
@@ -662,6 +703,7 @@ export default function ForwardPage() {
|
||||
inPort: form.inPort,
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizeSpeedId(form.speedId),
|
||||
};
|
||||
|
||||
res = await createForward(createData);
|
||||
@@ -1346,7 +1388,11 @@ export default function ForwardPage() {
|
||||
const aInx = a.inx ?? 0;
|
||||
const bInx = b.inx ?? 0;
|
||||
|
||||
return aInx - bInx;
|
||||
if (aInx !== bInx) {
|
||||
return aInx - bInx;
|
||||
}
|
||||
|
||||
return (a.id ?? 0) - (b.id ?? 0);
|
||||
});
|
||||
|
||||
// 如果数据库中没有排序信息,则使用本地存储的顺序
|
||||
@@ -1511,6 +1557,9 @@ export default function ForwardPage() {
|
||||
{forward.userName || "未知用户"}
|
||||
</span>
|
||||
</TableCell>
|
||||
<TableCell className="whitespace-nowrap font-semibold text-foreground">
|
||||
{forward.name}
|
||||
</TableCell>
|
||||
<TableCell className="whitespace-nowrap">
|
||||
<Chip
|
||||
className="border-none bg-secondary/10 px-2"
|
||||
@@ -1522,9 +1571,6 @@ export default function ForwardPage() {
|
||||
</span>
|
||||
</Chip>
|
||||
</TableCell>
|
||||
<TableCell className="whitespace-nowrap font-semibold text-foreground">
|
||||
{forward.name}
|
||||
</TableCell>
|
||||
<TableCell className="max-w-[220px]">
|
||||
<button
|
||||
className={`w-full truncate rounded-md bg-default-100/50 px-2.5 py-1.5 text-left font-mono text-xs font-medium text-default-700 transition-all ${
|
||||
@@ -2181,8 +2227,8 @@ export default function ForwardPage() {
|
||||
)}
|
||||
<TableColumn className="w-10 pl-4" />
|
||||
<TableColumn>用户</TableColumn>
|
||||
<TableColumn>隧道</TableColumn>
|
||||
<TableColumn>名称</TableColumn>
|
||||
<TableColumn>隧道</TableColumn>
|
||||
<TableColumn>入口</TableColumn>
|
||||
<TableColumn>目标</TableColumn>
|
||||
<TableColumn>策略</TableColumn>
|
||||
@@ -2301,6 +2347,38 @@ export default function ForwardPage() {
|
||||
}
|
||||
/>
|
||||
|
||||
<Select
|
||||
label="限速规则"
|
||||
placeholder="不限速"
|
||||
selectedKeys={
|
||||
selectedSpeedId !== null
|
||||
? [selectedSpeedId.toString()]
|
||||
: ["null"]
|
||||
}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as string;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
speedId:
|
||||
selectedKey === "null" ? null : Number(selectedKey),
|
||||
}));
|
||||
}}
|
||||
>
|
||||
<SelectItem key="null" textValue="不限速">
|
||||
不限速
|
||||
</SelectItem>
|
||||
{availableSpeedLimits.map((speedLimit) => (
|
||||
<SelectItem
|
||||
key={speedLimit.id.toString()}
|
||||
textValue={speedLimit.name}
|
||||
>
|
||||
{speedLimit.name}
|
||||
</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
|
||||
<Select
|
||||
description={
|
||||
isEdit
|
||||
|
||||
@@ -34,8 +34,8 @@ interface SpeedLimitRule {
|
||||
name: string;
|
||||
speed: number;
|
||||
status: number;
|
||||
tunnelId: number;
|
||||
tunnelName: string;
|
||||
tunnelId?: number | null;
|
||||
tunnelName?: string;
|
||||
createdTime: string;
|
||||
updatedTime: string;
|
||||
}
|
||||
@@ -139,9 +139,7 @@ export default function LimitPage() {
|
||||
newErrors.speed = "请输入有效的速度限制(≥1 Mbps)";
|
||||
}
|
||||
|
||||
if (!form.tunnelId) {
|
||||
newErrors.tunnelId = "请选择要绑定的隧道";
|
||||
}
|
||||
// tunnelId is optional - speed limits can be created without binding to a tunnel
|
||||
|
||||
setErrors(newErrors);
|
||||
|
||||
@@ -169,8 +167,8 @@ export default function LimitPage() {
|
||||
id: rule.id,
|
||||
name: rule.name,
|
||||
speed: rule.speed,
|
||||
tunnelId: rule.tunnelId,
|
||||
tunnelName: rule.tunnelName,
|
||||
tunnelId: rule.tunnelId ?? null,
|
||||
tunnelName: rule.tunnelName ?? "",
|
||||
status: rule.status,
|
||||
});
|
||||
setErrors({});
|
||||
@@ -219,6 +217,8 @@ export default function LimitPage() {
|
||||
const createData = { ...form };
|
||||
|
||||
delete createData.id;
|
||||
createData.tunnelId = null;
|
||||
createData.tunnelName = "";
|
||||
|
||||
res = await createSpeedLimit(createData);
|
||||
}
|
||||
@@ -393,9 +393,7 @@ export default function LimitPage() {
|
||||
{isEdit ? "编辑限速规则" : "新增限速规则"}
|
||||
</h2>
|
||||
<p className="text-small text-default-500">
|
||||
{isEdit
|
||||
? "修改现有限速规则的配置信息"
|
||||
: "创建新的限速规则并绑定到隧道"}
|
||||
{isEdit ? "修改现有限速规则的配置信息" : "创建新的限速规则"}
|
||||
</p>
|
||||
</ModalHeader>
|
||||
<ModalBody>
|
||||
@@ -435,43 +433,44 @@ export default function LimitPage() {
|
||||
}
|
||||
/>
|
||||
|
||||
<Select
|
||||
description={isEdit ? "编辑时无法修改绑定隧道" : undefined}
|
||||
errorMessage={errors.tunnelId}
|
||||
isDisabled={isEdit}
|
||||
isInvalid={!!errors.tunnelId}
|
||||
label="绑定隧道"
|
||||
placeholder="请选择要绑定的隧道"
|
||||
selectedKeys={
|
||||
form.tunnelId ? [form.tunnelId.toString()] : []
|
||||
}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as string;
|
||||
|
||||
if (selectedKey) {
|
||||
const selectedTunnel = tunnels.find(
|
||||
(tunnel) => tunnel.id === parseInt(selectedKey),
|
||||
);
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
tunnelId: parseInt(selectedKey),
|
||||
tunnelName: selectedTunnel?.name || "",
|
||||
}));
|
||||
} else {
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
tunnelId: null,
|
||||
tunnelName: "",
|
||||
}));
|
||||
{isEdit && (
|
||||
<Select
|
||||
description="仅编辑时可调整绑定隧道"
|
||||
errorMessage={errors.tunnelId}
|
||||
isInvalid={!!errors.tunnelId}
|
||||
label="绑定隧道"
|
||||
placeholder="可选择要绑定的隧道(可选)"
|
||||
selectedKeys={
|
||||
form.tunnelId ? [form.tunnelId.toString()] : []
|
||||
}
|
||||
}}
|
||||
>
|
||||
{tunnels.map((tunnel) => (
|
||||
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as string;
|
||||
|
||||
if (selectedKey) {
|
||||
const selectedTunnel = tunnels.find(
|
||||
(tunnel) => tunnel.id === parseInt(selectedKey),
|
||||
);
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
tunnelId: parseInt(selectedKey),
|
||||
tunnelName: selectedTunnel?.name || "",
|
||||
}));
|
||||
} else {
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
tunnelId: null,
|
||||
tunnelName: "",
|
||||
}));
|
||||
}
|
||||
}}
|
||||
>
|
||||
{tunnels.map((tunnel) => (
|
||||
<SelectItem key={tunnel.id}>{tunnel.name}</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
)}
|
||||
</div>
|
||||
</ModalBody>
|
||||
<ModalFooter>
|
||||
|
||||
@@ -366,6 +366,21 @@ export default function TunnelPage() {
|
||||
return form.chainNodes || [];
|
||||
};
|
||||
|
||||
const mergeOrderedNodes = (
|
||||
currentNodes: ChainTunnel[],
|
||||
selectedNodeIds: number[],
|
||||
buildDefault: (nodeId: number) => ChainTunnel,
|
||||
): ChainTunnel[] => {
|
||||
const selectedSet = new Set(selectedNodeIds);
|
||||
const kept = currentNodes.filter((node) => selectedSet.has(node.nodeId));
|
||||
const keptIds = new Set(kept.map((node) => node.nodeId));
|
||||
const added = selectedNodeIds
|
||||
.filter((nodeId) => !keptIds.has(nodeId))
|
||||
.map((nodeId) => buildDefault(nodeId));
|
||||
|
||||
return [...kept, ...added];
|
||||
};
|
||||
|
||||
// 提交表单
|
||||
const handleSubmit = async () => {
|
||||
if (!validateForm()) return;
|
||||
@@ -1241,15 +1256,10 @@ export default function TunnelPage() {
|
||||
const selectedIds = Array.from(keys).map((key) =>
|
||||
parseInt(key as string),
|
||||
);
|
||||
const newInNodeId: ChainTunnel[] = selectedIds.map(
|
||||
(nodeId) => {
|
||||
// 保留已有的端口配置
|
||||
const existing = form.inNodeId.find(
|
||||
(ct) => ct.nodeId === nodeId,
|
||||
);
|
||||
|
||||
return existing || { nodeId, chainType: 1 };
|
||||
},
|
||||
const newInNodeId = mergeOrderedNodes(
|
||||
form.inNodeId,
|
||||
selectedIds,
|
||||
(nodeId) => ({ nodeId, chainType: 1 }),
|
||||
);
|
||||
|
||||
setForm((prev) => ({ ...prev, inNodeId: newInNodeId }));
|
||||
@@ -1659,21 +1669,16 @@ export default function TunnelPage() {
|
||||
const realNodes = currentOutNodes.filter(
|
||||
(ct) => ct.nodeId !== -1,
|
||||
);
|
||||
const newOutNodeId: ChainTunnel[] =
|
||||
selectedIds.map((nodeId) => {
|
||||
const existing = realNodes.find(
|
||||
(ct) => ct.nodeId === nodeId,
|
||||
);
|
||||
|
||||
return (
|
||||
existing || {
|
||||
nodeId,
|
||||
chainType: 3,
|
||||
protocol,
|
||||
strategy,
|
||||
}
|
||||
);
|
||||
});
|
||||
const newOutNodeId = mergeOrderedNodes(
|
||||
realNodes,
|
||||
selectedIds,
|
||||
(nodeId) => ({
|
||||
nodeId,
|
||||
chainType: 3,
|
||||
protocol,
|
||||
strategy,
|
||||
}),
|
||||
);
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { useState, useEffect } from "react";
|
||||
import { useState, useEffect, useMemo } from "react";
|
||||
import toast from "react-hot-toast";
|
||||
import { parseDate } from "@internationalized/date";
|
||||
|
||||
@@ -219,6 +219,22 @@ export default function UserPage() {
|
||||
const [speedLimits, setSpeedLimits] = useState<SpeedLimit[]>([]);
|
||||
const [userGroups, setUserGroups] = useState<UserGroup[]>([]);
|
||||
|
||||
const noLimitSpeedLimitIds = useMemo(() => {
|
||||
return new Set(
|
||||
speedLimits
|
||||
.filter((speedLimit) => speedLimit.name.trim() === "不限速")
|
||||
.map((speedLimit) => speedLimit.id),
|
||||
);
|
||||
}, [speedLimits]);
|
||||
|
||||
const normalizeSpeedId = (speedId?: number | null): number | null => {
|
||||
if (speedId === null || speedId === undefined) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return noLimitSpeedLimitIds.has(speedId) ? null : speedId;
|
||||
};
|
||||
|
||||
// 生命周期
|
||||
useEffect(() => {
|
||||
loadUsers();
|
||||
@@ -432,7 +448,10 @@ export default function UserPage() {
|
||||
try {
|
||||
const tunnelsToAssign: TunnelAssignItem[] = Array.from(
|
||||
batchTunnelSelections.entries(),
|
||||
).map(([tunnelId, speedId]) => ({ tunnelId, speedId }));
|
||||
).map(([tunnelId, speedId]) => ({
|
||||
tunnelId,
|
||||
speedId: normalizeSpeedId(speedId),
|
||||
}));
|
||||
|
||||
const response = await batchAssignUserTunnel({
|
||||
userId: currentUser.id,
|
||||
@@ -456,6 +475,7 @@ export default function UserPage() {
|
||||
const handleEditTunnel = (userTunnel: UserTunnel) => {
|
||||
setEditTunnelForm({
|
||||
...userTunnel,
|
||||
speedId: normalizeSpeedId(userTunnel.speedId),
|
||||
expTime: userTunnel.expTime,
|
||||
});
|
||||
onEditTunnelModalOpen();
|
||||
@@ -472,7 +492,7 @@ export default function UserPage() {
|
||||
num: editTunnelForm.num,
|
||||
expTime: editTunnelForm.expTime,
|
||||
flowResetTime: editTunnelForm.flowResetTime,
|
||||
speedId: editTunnelForm.speedId,
|
||||
speedId: normalizeSpeedId(editTunnelForm.speedId),
|
||||
status: editTunnelForm.status,
|
||||
});
|
||||
|
||||
@@ -583,13 +603,17 @@ export default function UserPage() {
|
||||
};
|
||||
|
||||
const editAvailableSpeedLimits = speedLimits.filter(
|
||||
(speedLimit) => speedLimit.tunnelId === editTunnelForm?.tunnelId,
|
||||
(speedLimit) => !noLimitSpeedLimitIds.has(speedLimit.id),
|
||||
);
|
||||
|
||||
const getSpeedLimitsForTunnel = (tunnelId: number) => {
|
||||
return speedLimits.filter((sl) => sl.tunnelId === tunnelId);
|
||||
const getSpeedLimitsForTunnel = (_tunnelId: number) => {
|
||||
return speedLimits.filter(
|
||||
(speedLimit) => !noLimitSpeedLimitIds.has(speedLimit.id),
|
||||
);
|
||||
};
|
||||
|
||||
const editTunnelSelectedSpeedId = normalizeSpeedId(editTunnelForm?.speedId);
|
||||
|
||||
const toggleTunnelSelection = (tunnelId: number) => {
|
||||
setBatchTunnelSelections((prev) => {
|
||||
const newMap = new Map(prev);
|
||||
@@ -1406,8 +1430,8 @@ export default function UserPage() {
|
||||
<Select
|
||||
label="限速规则"
|
||||
selectedKeys={
|
||||
editTunnelForm.speedId
|
||||
? [editTunnelForm.speedId.toString()]
|
||||
editTunnelSelectedSpeedId !== null
|
||||
? [editTunnelSelectedSpeedId.toString()]
|
||||
: ["null"]
|
||||
}
|
||||
onSelectionChange={(keys) => {
|
||||
|
||||
@@ -88,7 +88,8 @@ export interface Tunnel {
|
||||
export interface SpeedLimit {
|
||||
id: number;
|
||||
name: string;
|
||||
tunnelId: number;
|
||||
tunnelId?: number | null;
|
||||
speed?: number;
|
||||
uploadSpeed: number;
|
||||
downloadSpeed: number;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user