Compare commits

...

52 Commits

Author SHA1 Message Date
sagitchu 3373e5ade9 fix: preserve shared limiters with per-IP rules 2026-04-28 00:18:32 +08:00
sagitchu bd27b94909 fix: omit per-IP speed payload for users 2026-04-27 23:34:35 +08:00
sagitchu a5a500bc0f feat: add per-IP limit controls 2026-04-27 23:17:02 +08:00
sagitchu edfe2a2372 fix: preserve udp packet semantics with limiters 2026-04-27 23:13:16 +08:00
sagitchu 7a9ba8bd81 fix: apply per-client limits to udp listener 2026-04-27 23:05:31 +08:00
sagitchu 46394388b1 feat: sync per-IP runtime limiters 2026-04-27 22:32:35 +08:00
sagitchu dec337d46b fix: enforce per-IP speed update permissions 2026-04-27 22:26:34 +08:00
sagitchu 9e8d27d98e feat: expose per-IP forward limit fields 2026-04-27 22:20:26 +08:00
sagitchu a2000e4d98 fix: preserve per-IP forward limits during rollback 2026-04-27 22:16:38 +08:00
sagitchu 2b76a9f0be feat: persist per-IP forward limits 2026-04-27 22:07:04 +08:00
sagitchu 54d7dfb7c9 docs: design per-IP rule limits 2026-04-27 17:24:25 +08:00
sagit 9f19d5fe15 fix: send valid announcement update payload 2026-04-27 15:34:47 +08:00
sagit 2ca3849917 fix: apply proxy protocol and max connection settings 2026-04-27 10:54:25 +08:00
sagit 58d2e89147 fix: reduce reconnect redeploy and metrics load (#476)
* fix: reduce reconnect redeploy and metrics load

Throttle node-online redeploy retries and lower the agent metric cadence so brief reconnect churn no longer fans out into repeated runtime syncs and backend connection pressure.

* docs: add follow-up implementation design notes

Document the planned flow upload batching work and the local remote-address toggle so the next changesets can implement them against an agreed design.
2026-04-26 23:35:35 +08:00
sagit 87a1a34ad5 refactor: batch flow upload processing (#474)
* test: cover flow upload batch semantics

* refactor: batch flow upload persistence

* refactor: batch flow upload processing

* test: harden flow upload batch regression coverage
2026-04-26 20:45:48 +08:00
sagit a625884d61 fix: remove unwanted hover underline from forward advanced settings (#473)
* fix: remove forward advanced settings hover underline

* fix: clear frontend lint warnings
2026-04-26 19:12:38 +08:00
sagit 799bb66fe5 feat: add local remote address toggle (#472) 2026-04-26 16:30:05 +08:00
sagit 3f374df724 fix: 保留用户/规则高级设置并增强节点运行时恢复 (#468)
* fix: 用户编辑时 maxConn 字段未正确回填

- User 接口添加 maxConn 类型定义
- handleEdit 中回填 maxConn 值
- normalizeUserItem 添加 maxConn 字段处理

解决编辑用户时最大连接数显示为 0 的问题

* fix: 保留转发设置并增强节点运行时恢复
2026-04-25 18:12:13 +08:00
sagit 9b923a2d0b fix: remove duplicate maxConn input in user form (#467) 2026-04-24 21:56:46 +08:00
sagit d9f28f53c7 style: move advanced settings to bottom and optimize UI (#466) 2026-04-24 19:10:36 +08:00
sagit 1d08a1ccfc fix: ensure page refresh to login page upon logout (#465) 2026-04-23 20:43:27 +08:00
sagit db25ba2cbe style: move maxConn and speed limit inputs to advanced settings accordion (#464) 2026-04-23 20:29:52 +08:00
sagit a498067261 fix: properly initialize maxConn in forward rules (#463)
* fix: initialize maxConn in forward creation and edit forms

* fix: add maxConn to Forward interface
2026-04-23 20:23:18 +08:00
sagit b5922dccf2 fix: include maxConn when submitting forward forms (#462)
* fix: frontend missing maxConn setting during forward creation

* docs: check off plan
2026-04-23 20:11:20 +08:00
sagit c259645227 feat: implement user and forward max connection limit (#461)
* docs: add implementation plan for max conn limit

* feat: add CLimiters support for websocket reporter

* feat: add max_conn field to user and forward models

* feat: implement max conn limiter dispatching

* feat: add maxConn to user and forward CRUD API

* feat: add max conn UI to user management and forward rules

* fix: load MaxConn in get forward record and add e2e contract test for max conn limit
2026-04-23 17:35:39 +08:00
sagitchu bdc2c4ecbb fix(backend): dynamically start/stop tunnel quality prober based on config 2026-04-23 14:14:24 +08:00
sagitchu a070d0f4d3 fix(backend): allow setting reserved ip addresses as target
Only forbid internal networks (loopback and private ips), removing restrictions on reserved addresses like multicast or unspecified.
2026-04-22 19:00:21 +08:00
sagitchu eecdd62d3a perf(backend): optimize kcp tunnel parameters for high throughput and low latency 2026-04-22 16:29:02 +08:00
sagitchu e6d3b847bb fix(backend): properly identify and clean up orphaned tunnel_%d services 2026-04-22 16:18:57 +08:00
sagitchu 0b49cd720f fix(backend): use proper protocol for chain hop diagnosis and tcp for external targets 2026-04-22 14:02:59 +08:00
sagitchu 4f488ae7ef fix(backend): properly implement tunnel_%d cleanup that was lost during revert 2026-04-22 11:32:49 +08:00
sagitchu 2b2b417f91 fix(gost): ensure JSON floats are correctly converted in GetInt and remove accidental binary 2026-04-22 11:17:59 +08:00
sagitchu 01b4c3e3eb fix: add workbox-window as explicit dependency to resolve pnpm strict resolution in CI 2026-04-22 11:05:37 +08:00
sagitchu e7c967df00 chore: migrate frontend package manager from npm to pnpm 2026-04-22 11:02:43 +08:00
sagitchu efaf920e51 fix: ensure rollbackTunnelRuntime includes tunnel_%d in cleanup 2026-04-22 09:51:35 +08:00
sagitchu 9aff669c0e fix: ensure tunnel protocol and KCP config are correctly processed and cleaned up 2026-04-22 09:41:41 +08:00
sagitchu 6e60f5cfbd fix: add visually hidden DialogTitle to resolve a11y warning and fix modal background 2026-04-22 09:21:47 +08:00
sagitchu a03c320b89 fix(gost): prevent panic in service parsing with empty md 2026-04-21 22:11:22 +08:00
sagitchu 61dba0ae57 fix: remove trailing brace in tunnel.tsx to resolve CI error 2026-04-21 20:20:19 +08:00
sagitchu f3d6366471 fix: increase kcp tunnel bandwidth limit and fix modal backgrounds 2026-04-21 20:16:58 +08:00
sagitchu c431d79403 fix: correct tunnel protocol handling for KCP cleanup and diagnosis
- Fix KCP tunnel not reclaimed after deletion: service name was
  hardcoded to {id}_tls but services were created as {id}_kcp,
  causing DeleteService to never find the actual KCP service.
  Now reads tunnel.Protocol from DB and derives correct name.

- Fix KCP diagnosis using TCP ping instead of UDP ping:
  tunnel.Protocol was always hardcoded to 'tls' at creation,
  so isUDPBasedProtocol() never matched kcp tunnels. Now
  stores the actual protocol from entry node configuration.

- Fix addTunnelServiceOnNode to extract service name from
  serviceData instead of hardcoding _tls suffix.

- Fix rollbackTunnelRuntime to accept protocol parameter
  so retry cleanup uses correct service name.

- UpdateTunnelTx now persists protocol on tunnel updates.
2026-04-21 18:48:11 +08:00
sagitchu 5107f59d94 fix: tolerate offline nodes when controlling forward services and editing tunnels
- controlForwardServices: skip offline nodes instead of failing entire operation,
  so forward pause/resume/delete works when some entry nodes are offline
- onNodeOnline: always sync forward state on node reconnect (not just post-upgrade),
  so forwards that changed status while a node was offline get synced
- add ListForwardIDsByNode repo method to sync all forwards (including paused)
- tunnel edit UI: allow deselecting already-selected offline nodes in
  entry/chain/exit selectors, matching the backend's existing tolerance
2026-04-21 16:53:47 +08:00
sagitchu 431613cb6a fix(frontend): listen to session changes so logout redirects to login page 2026-04-21 16:24:46 +08:00
sagitchu a6b218f3ee fix(frontend): increase sidebar background opacity to reduce grayness 2026-04-21 16:12:28 +08:00
sagitchu d5d26d9cf9 fix(frontend): hide scrollbars on sidebar, modals, and secondary lists; fix diagnosis modal bg 2026-04-21 16:00:06 +08:00
sagitchu c1bc795674 fix(kcp): enable congestion control, add FEC, and fix remote node UDP diagnosis
- KCP: switch from fast3 to fast2 mode, enable FEC (10/3), enable
  congestion control (nc=0) for automatic rate adaptation
- go-gost/x: support kcp.nc metadata to override mode-initialized
  NoCongestion value in dialer and listener Init()
- Diagnosis: add udpPingViaRemoteNode, fix pingViaRemoteNode to
  dispatch based on protocol, add Protocol field to federation
  diagnose request structs
2026-04-21 15:41:32 +08:00
sagitchu 30e1473f06 fix(traffic): fix flow counter inflation from TOCTOU race in agent traffic reporter
Replace read-then-subtract pattern in collectAndReport with atomic
swap-to-zero to eliminate race where AddTraffic increments counters
between snapshot and clearReportedTraffic, causing residual traffic
to accumulate indefinitely and inflate user flow counters.

Also add defensive check in processFlowItem to skip AddFlow when
forward no longer exists, and send DeleteService to clean up orphaned
agent services.
2026-04-21 14:32:08 +08:00
sagitchu eaf16bf17b fix(frontend): enable horizontal scroll on all table lists for mobile
Change overflow-hidden to overflow-auto in table wrapper classNames
across all list pages (rules, tunnels, nodes, limits, groups, users,
monitor views). Also add table-fixed min-w to forward all-rules table
and fix unescaped entities in dashboard.
2026-04-21 14:18:44 +08:00
sagitchu e995d70be7 feat(diagnosis): add protocol-aware connectivity diagnosis for tunnel chains
- Pass tunnel protocol through diagnosis work items
- Support protocol-specific ping (TCP/UDP/KCP) via remote nodes
- Add KCP probe support in websocket_reporter for chain hop testing
2026-04-21 14:00:02 +08:00
sagitchu a968a10792 fix(frontend): fix dropdown padding, hide scrollbar, login input transparency
- Increase dropdown menu padding from p-1 to p-1.5
- Hide scrollbar in dropdown-menu-content and select listbox
- Change login page input bg from white/50 to white/10 for transparency
2026-04-21 12:41:49 +08:00
sagitchu ea156c33bc style(frontend): sync remaining glass UI theme files and cyberpunk theme 2026-04-21 12:40:00 +08:00
sagitchu 29407c90b6 perf(kcp): optimize tunnel transport defaults for high throughput
- Set KCP mode to fast3 (NoDelay:1, Interval:10ms) for lower latency
- Double SndWnd/RcvWnd from 1024 to 2048 for higher BDP
- Disable FEC (datashard:0, parityshard:0) to eliminate 30% overhead
- Disable compression for tunnel transport
- Increase relay mux MaxStreamBuffer to 2MB for better UDP throughput
- Add kcp.datashard/kcp.parityshard metadata keys support

Before: 210Mbps TCP / 35Mbps UDP (93% loss)
After:  should approach direct-connection speed (~400Mbps+)
2026-04-21 11:33:41 +08:00
116 changed files with 19173 additions and 2073 deletions
+5 -2
View File
@@ -21,11 +21,14 @@ jobs:
with:
node-version: '20.19.0'
- name: Install pnpm
run: npm install -g pnpm
- name: Install dependencies
run: npm install --legacy-peer-deps
run: pnpm install --frozen-lockfile
- name: Build
run: npm run build
run: pnpm run build
backend:
name: Build Go Backend
+4 -4
View File
@@ -22,10 +22,10 @@ FLVX — traffic forwarding panel: Go admin API + Vite/React UI + Go agent.
(cd go-backend && go test ./...)
# Frontend
(cd vite-frontend && npm install --legacy-peer-deps)
(cd vite-frontend && npm run dev) # host 0.0.0.0:3000
(cd vite-frontend && npm run build) # tsc && vite build
(cd vite-frontend && npm run lint) # eslint --fix (no typecheck command)
(cd vite-frontend && pnpm install)
(cd vite-frontend && pnpm run dev) # host 0.0.0.0:3000
(cd vite-frontend && pnpm run build) # tsc && vite build
(cd vite-frontend && pnpm run lint) # eslint --fix (no typecheck command)
# Agent
(cd go-gost && go run .)
@@ -0,0 +1,482 @@
# 最大连接数限制实现计划
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`[x]`) syntax for tracking.
**Goal:** 在 FLVX 中实现基于用户的全局最大连接数限制和基于单条规则的独立最大连接数限制功能。前端输入框为 0 或空时表示不限制。
**Architecture:** 采用“覆盖逻辑”(方案二)。
1. 数据库层面:在 `user` 和 `forward` 表中各增加一个整型字段 `max_conn`,默认值为 0(表示不限制)。
2. 后端接口层面:提供 API 更新该字段,在组装下发给 GOST 的配置时,判断规则的 `max_conn` 是否大于 0:
- 如果规则 `max_conn > 0`,则为此规则动态生成一个唯一的连接限制器配置,并在下发服务的 `climiter` 字段中引用该限制器。
- 如果规则 `max_conn == 0`,则检查该规则所属用户的 `max_conn`。
- 如果用户 `max_conn > 0`,则引用以用户维度的连接限制器配置(如 `user_conn_limit_<user_id>`)。
- 否则不下发 `climiter`。
3. 后端服务控制平面:需要在下发服务前,将需要的连接限制器(Rule 或 User 维度)推送到节点上。
- **重要发现:** 当前 `go-gost` 的 WebSocket Reporter (`go-gost/x/socket/websocket_reporter.go`) 仅支持 `TrafficLimiter` 的动态增删(如 `AddLimiters` 等),**不支持** `ConnLimiter`(即 `CLimiters`)。
- **计划修改:** 我们需要先在 `go-gost` 侧(`go-gost/x/socket`)添加针对 `CLimiters` 的 WebSocket 指令(`AddCLimiters`, `UpdateCLimiters`, `DeleteCLimiters`)以及对应的处理函数(参考 `AddLimiters` 等的实现,调用现有的针对 `ConnLimiterRegistry` 的相关接口和配置存储逻辑,具体需要实现类似 `createLimiter` 到 `createConnLimiter` 的逻辑)。
- 完成底层修改后,`go-backend` 再通过这些新增加的 WebSocket 指令,在 `ensureLimiterOnNode` 时下发最大连接数限制规则。
4. 前端层面:在用户管理和规则管理页面增加输入框组件。
**Tech Stack:** Go, GORM, SQLite/PostgreSQL, React, Vite, TypeScript, TailwindCSS.
---
### Task 1: 扩展 go-gost WebSocket 接口以支持 CLimiters
**Files:**
- Modify: `go-gost/x/socket/limiter.go`
- Modify: `go-gost/x/socket/websocket_reporter.go`
[x] **Step 1: 实现 `createConnLimiter` 等功能**
在 `go-gost/x/socket/limiter.go` 中参考现有 `createLimiter` 添加对 `CLimiters` 的支持:
```go
func createConnLimiter(req createLimiterRequest) error {
name := strings.TrimSpace(req.Data.Name)
if name == "" {
return errors.New("limiter name is required")
}
req.Data.Name = name
if registry.ConnLimiterRegistry().IsRegistered(name) {
return errors.New("conn limiter " + name + " already exists")
}
v := parser.ParseConnLimiter(&req.Data)
if err := registry.ConnLimiterRegistry().Register(name, v); err != nil {
return errors.New("conn limiter " + name + " already exists")
}
if c := config.Global(); c != nil {
c.CLimiters = append(c.CLimiters, &req.Data)
}
return nil
}
func updateConnLimiter(req updateLimiterRequest) error {
name := strings.TrimSpace(req.Limiter)
req.Data.Name = name
if registry.ConnLimiterRegistry().IsRegistered(name) {
registry.ConnLimiterRegistry().Unregister(name)
}
v := parser.ParseConnLimiter(&req.Data)
if err := registry.ConnLimiterRegistry().Register(name, v); err != nil {
return errors.New("conn limiter " + name + " already exists")
}
if c := config.Global(); c != nil {
for i := range c.CLimiters {
if c.CLimiters[i].Name == name {
c.CLimiters[i] = &req.Data
return nil
}
}
c.CLimiters = append(c.CLimiters, &req.Data)
}
return nil
}
func deleteConnLimiter(req deleteLimiterRequest) error {
name := strings.TrimSpace(req.Limiter)
if registry.ConnLimiterRegistry().IsRegistered(name) {
registry.ConnLimiterRegistry().Unregister(name)
}
if c := config.Global(); c != nil {
limiteres := c.CLimiters
c.CLimiters = nil
for _, s := range limiteres {
if s.Name == name {
continue
}
c.CLimiters = append(c.CLimiters, s)
}
}
return nil
}
```
[x] **Step 2: 在 `WebSocketReporter` 注册命令**
在 `go-gost/x/socket/websocket_reporter.go` 的 `ProcessCommand` 中添加 case:
```go
case "AddCLimiters":
err = w.handleAddCLimiter(cmd.Data)
response.Type = "AddCLimitersResponse"
needSaveConfig = true
case "UpdateCLimiters":
err = w.handleUpdateCLimiter(cmd.Data)
response.Type = "UpdateCLimitersResponse"
needSaveConfig = true
case "DeleteCLimiters":
err = w.handleDeleteCLimiter(cmd.Data)
response.Type = "DeleteCLimitersResponse"
needSaveConfig = true
```
[x] **Step 3: 实现 Handler 方法**
在 `go-gost/x/socket/websocket_reporter.go` 中添加:
```go
func (w *WebSocketReporter) handleAddCLimiter(data interface{}) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("序列化数据失败: %v", err)
}
var limiterConfig config.LimiterConfig
if err := json.Unmarshal(jsonData, &limiterConfig); err != nil {
return fmt.Errorf("解析限流器配置失败: %v", err)
}
req := createLimiterRequest{Data: limiterConfig}
return createConnLimiter(req)
}
func (w *WebSocketReporter) handleUpdateCLimiter(data interface{}) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("序列化数据失败: %v", err)
}
var updateReq struct {
Limiter string `json:"limiter"`
Data config.LimiterConfig `json:"data"`
}
if err := json.Unmarshal(jsonData, &updateReq); err != nil {
var limiterConfig config.LimiterConfig
if err := json.Unmarshal(jsonData, &limiterConfig); err != nil {
return fmt.Errorf("解析更新请求失败: %v", err)
}
updateReq.Limiter = limiterConfig.Name
updateReq.Data = limiterConfig
}
req := updateLimiterRequest{
Limiter: updateReq.Limiter,
Data: updateReq.Data,
}
return updateConnLimiter(req)
}
func (w *WebSocketReporter) handleDeleteCLimiter(data interface{}) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("序列化数据失败: %v", err)
}
var deleteReq deleteLimiterRequest
if err := json.Unmarshal(jsonData, &deleteReq); err != nil {
var limiterName string
if err := json.Unmarshal(jsonData, &limiterName); err != nil {
return fmt.Errorf("解析删除请求失败: %v", err)
}
deleteReq.Limiter = limiterName
}
return deleteConnLimiter(deleteReq)
}
```
[x] **Step 4: Commit**
```bash
cd go-gost
git add x/socket/limiter.go x/socket/websocket_reporter.go
git commit -m "feat: add CLimiters support for websocket reporter"
cd ..
```
---
### Task 2: 数据库迁移与模型更新
**Files:**
- Modify: `go-backend/internal/store/model/model.go`
- Modify: `go-backend/internal/store/repo/repository.go`
[x] **Step 1: 更新数据库模型**
在 `go-backend/internal/store/model/model.go` 的 `User` 和 `Forward` 结构体中添加 `MaxConn` 字段。
```go
// 在 User 结构体中
type User struct {
// ...
MaxConn int `gorm:"column:max_conn;not null;default:0"`
// ...
}
// 在 Forward 结构体中
type Forward struct {
// ...
MaxConn int `gorm:"column:max_conn;not null;default:0"`
// ...
}
```
[x] **Step 2: 编写数据库迁移**
在 `go-backend/internal/store/repo/repository.go` 的 `AutoMigrate` 逻辑前(如果有自定义迁移)或利用 gorm 自动迁移机制,由于这是 autoMigrate,添加字段只要 `db.AutoMigrate(&model.User{}, &model.Forward{})` 被调用就能自动加上。确认已执行迁移。由于 `FLVX` 通常会自动执行迁移,只需修改模型即可。我们需要处理默认值,由于使用了 `default:0`,GORM 会处理新增字段的默认值,但为了安全起见,如果在旧环境中,可能直接 alter table。
```go
// 无需手动编写 SQL,依赖现有的 gorm AutoMigrate 即可。
```
[x] **Step 3: 运行并验证迁移通过**
Run: `make build` (在 go-backend 中),或者运行一个相关的存储单元测试。
[x] **Step 4: Commit**
```bash
cd go-backend
git add internal/store/model/model.go
git commit -m "feat: add max_conn field to user and forward models"
cd ..
```
---
### Task 3: 后端控制平面 - 连接数限制器的组装与下发
**Files:**
- Modify: `go-backend/internal/http/handler/control_plane.go`
- Modify: `go-backend/internal/store/repo/repository_control.go`
[x] **Step 1: 更新存储层以获取 User 的 MaxConn**
在 `go-backend/internal/store/repo/repository_control.go` 中:
需要一个方法获取 User,或者如果已经有,确保可以拿到 `MaxConn`。
[x] **Step 2: 编写下发 CLimiter 到节点的辅助函数**
在 `go-backend/internal/http/handler/control_plane.go`,参考 `ensureLimiterOnNode` 和 `upsertLimiterOnNode`:
```go
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error {
limitStr := fmt.Sprintf("$ %d", maxConn)
payload := map[string]interface{}{
"name": limiterName,
"limits": []string{limitStr},
}
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
if !isAlreadyExistsMessage(err.Error()) {
return fmt.Errorf("连接限制器下发失败: %w", err)
}
updatePayload := map[string]interface{}{
"limiter": limiterName,
"data": payload,
}
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
}
}
return nil
}
```
[x] **Step 3: 更新组装配置逻辑以绑定 `climiter`**
在 `control_plane.go` 的 `syncForwardServicesWithWarnings` 及其辅助函数 `buildForwardServiceConfigs` 附近:
修改 `buildForwardServiceConfigs` 的签名,传入 `maxConn int` 和对应的 `cLimiterName string`。
```go
func buildForwardServiceConfigs(baseName string, forward *model.Forward, tunnel *model.Tunnel, node *model.Node, port int, bindIP string, limiterID *int64, cLimiterName string) []map[string]interface{} {
// ... 现有逻辑
// 在服务配置生成的部分增加:
if cLimiterName != "" {
service["climiter"] = cLimiterName
}
// ...
}
```
[x] **Step 4: 在转发服务同步主流程中决定并下发 `climiter`**
在 `syncForwardServicesWithWarnings` (可能在多个重载/处理入口处,如 `ensureForwardServices`),查出转发所属 user 的 `MaxConn`,以及转发本身的 `MaxConn`。
```go
// 获取 User
user, err := h.repo.GetUser(forward.UserID)
if err != nil {
return nil, err
}
var cLimiterName string
var maxConnToSet int
if forward.MaxConn > 0 {
maxConnToSet = forward.MaxConn
cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID)
} else if user != nil && user.MaxConn > 0 {
maxConnToSet = user.MaxConn
cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID)
}
if cLimiterName != "" {
for _, fp := range ports {
if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil {
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
}
}
}
// 传递给 buildForwardServiceConfigs
// ...
```
*(注意:需要确保更新涉及 `buildForwardServiceConfigs` 的所有调用点)*
[x] **Step 5: Commit**
```bash
cd go-backend
git add internal/http/handler/control_plane.go internal/store/repo/repository_control.go
git commit -m "feat: implement max conn limiter dispatching"
cd ..
```
---
### Task 4: 后端接口 - 用户和规则的 CRUD 支持
**Files:**
- Modify: `go-backend/internal/http/handler/admin_user.go`
- Modify: `go-backend/internal/http/handler/forward.go`
[x] **Step 1: 用户接口更新**
在 `go-backend/internal/http/handler/admin_user.go`,修改用户创建和更新请求的结构体(如果有),接收 `MaxConn`,并在保存到数据库时赋值。
```go
type CreateUserReq struct {
// ...
MaxConn *int `json:"maxConn"`
}
// 接收后:
if req.MaxConn != nil {
user.MaxConn = *req.MaxConn
}
```
在获取用户列表时,确保 `MaxConn` 返回给前端。
[x] **Step 2: 规则接口更新**
在 `go-backend/internal/http/handler/forward.go` 中,更新 `CreateForwardReq` 和 `UpdateForwardReq` 结构体,增加 `MaxConn`,并在创建/更新 Forward 时保存到数据库。
如果转发规则的 `MaxConn` 或相关信息改变,触发节点上的规则重载(重新下发服务)。这一步由于更改了数据库,复用现有的 `syncForwardServices` 就会带上最新的配置。
[x] **Step 3: 测试接口**
Run: 可以启动后使用 curl 测试。
[x] **Step 4: Commit**
```bash
cd go-backend
git add internal/http/handler/admin_user.go internal/http/handler/forward.go
git commit -m "feat: add maxConn to user and forward CRUD API"
cd ..
```
---
### Task 5: 前端 - 用户管理页面集成
**Files:**
- Modify: `vite-frontend/src/api/types.ts`
- Modify: `vite-frontend/src/api/index.ts`
- Modify: `vite-frontend/src/pages/users.tsx` (或者对应的用户管理页面文件)
[x] **Step 1: 类型更新**
在 `vite-frontend/src/api/types.ts` 中:
为 `UserApiItem` 和相关的 mutation payload 增加 `maxConn?: number` 属性。
[x] **Step 2: UI 修改**
在用户创建/编辑弹窗中,增加“最大连接数”输入框:
(假设使用 `@nextui-org/react` 的 `Input`)
```tsx
<Input
type="number"
label="最大并发连接数"
placeholder="0 或空表示不限制"
value={formData.maxConn === 0 ? "" : String(formData.maxConn || "")}
onValueChange={(val) => {
const num = parseInt(val, 10);
setFormData({ ...formData, maxConn: isNaN(num) ? 0 : num });
}}
/>
```
并在用户的表格列中展示 `最大连接数`(值为 0 显示“不限制”)。
[x] **Step 3: 运行 Vite 进行验证**
[x] **Step 4: Commit**
```bash
cd vite-frontend
git add src/api/types.ts src/api/index.ts src/pages/users.tsx
git commit -m "feat: add max conn UI to user management"
cd ..
```
---
### Task 6: 前端 - 转发规则页面集成
**Files:**
- Modify: `vite-frontend/src/pages/forward.tsx`
[x] **Step 1: 类型更新**
在 `api/types.ts` 中 `ForwardMutationPayload` 和 `ForwardApiItem` 中增加 `maxConn?: number`。
[x] **Step 2: UI 修改**
在 `vite-frontend/src/pages/forward.tsx` 的创建/编辑规则弹窗(在 "规则限速" 附近)增加“最大连接数”输入框:
```tsx
<Input
type="number"
label="最大并发连接数"
placeholder="0 或空表示不限制"
value={formData.maxConn === 0 ? "" : String(formData.maxConn || "")}
onValueChange={(val) => {
const num = parseInt(val, 10);
setFormData({ ...formData, maxConn: isNaN(num) ? 0 : num });
}}
description="此设置优先于用户的全局连接数限制。0 表示不限制(或使用用户的全局限制)。"
/>
```
如果是在列表/卡片中展示,可以增加一个小标签或者 Tooltip 显示其最大连接数设置。
[x] **Step 3: 验证**
在前端验证该功能能正确读写规则的连接限制字段。
[x] **Step 4: Commit**
```bash
cd vite-frontend
git add src/pages/forward.tsx src/api/types.ts
git commit -m "feat: add max conn UI to forward rules"
cd ..
```
@@ -0,0 +1,171 @@
# Allow Local Remote Address Implementation Plan
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
**Goal:** Add a global settings toggle that allows non-admin forward rules to target local/private addresses when explicitly enabled.
**Architecture:** Keep the existing remote-address safety validator as the default path for non-admin rule changes, but gate its use behind a single backend config lookup in forward create/update handlers. Surface the toggle through the existing `vite_config` settings page and prove behavior with backend contract tests first.
**Tech Stack:** Go `net/http` + GORM backend, React + TypeScript frontend settings page, Go contract tests.
---
### Task 1: Backend Contract Coverage
**Files:**
- Modify: `go-backend/tests/contract/forward_contract_test.go`
- [ ] **Step 1: Write the failing tests**
Add contract tests that prove the desired behavior:
```go
t.Run("local remote address is rejected when toggle is off", func(t *testing.T) {
createPayload := map[string]interface{}{
"name": "deny-local-remote",
"tunnelId": tunnelID,
"remoteAddr": "127.0.0.1:8080",
"strategy": "fifo",
}
createBody, _ := json.Marshal(createPayload)
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
createReq.Header.Set("Authorization", adminToken)
createReq.Header.Set("Content-Type", "application/json")
createRes := httptest.NewRecorder()
router.ServeHTTP(createRes, createReq)
var out response.R
_ = json.NewDecoder(createRes.Body).Decode(&out)
if out.Code == 0 {
t.Fatalf("expected local remote address to be rejected when toggle is off")
}
})
t.Run("local remote address is allowed when toggle is on", func(t *testing.T) {
if err := repo.DB().Exec(`
INSERT INTO vite_config(name, value, time)
VALUES(?, ?, ?)
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
`, "allow_local_remote_addr", "1", time.Now().UnixMilli()).Error; err != nil {
t.Fatalf("enable allow_local_remote_addr: %v", err)
}
createPayload := map[string]interface{}{
"name": "allow-local-remote",
"tunnelId": tunnelID,
"remoteAddr": "127.0.0.1:8080",
"strategy": "fifo",
}
createBody, _ := json.Marshal(createPayload)
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
createReq.Header.Set("Authorization", adminToken)
createReq.Header.Set("Content-Type", "application/json")
createRes := httptest.NewRecorder()
router.ServeHTTP(createRes, createReq)
assertCode(t, createRes, 0)
})
```
- [ ] **Step 2: Run tests to verify they fail**
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
Expected: FAIL because backend still rejects local/private addresses unconditionally.
- [ ] **Step 3: Commit**
Do not commit yet; combine with Task 2 after implementation passes.
### Task 2: Backend Toggle Implementation
**Files:**
- Modify: `go-backend/internal/http/handler/mutations.go`
- [ ] **Step 1: Add a tiny config helper**
Add a helper near other handler helpers:
```go
func (h *Handler) allowLocalRemoteAddr() bool {
if h == nil || h.repo == nil {
return false
}
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
if err != nil || cfg == nil {
return false
}
return strings.TrimSpace(cfg.Value) == "1"
}
```
- [ ] **Step 2: Gate create/update validation behind the helper**
Replace the unconditional checks with:
```go
if !h.allowLocalRemoteAddr() {
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
response.WriteJSON(w, response.Err(403, err.Error()))
return
}
}
```
- [ ] **Step 3: Run contract tests to verify they pass**
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
Expected: PASS
- [ ] **Step 4: Run full backend tests**
Run: `go test ./...`
Expected: PASS
### Task 3: Settings Page Toggle
**Files:**
- Modify: `vite-frontend/src/pages/config.tsx`
- [ ] **Step 1: Add the config item to the settings schema**
Add a switch-style item for `allow_local_remote_addr` with warning copy about reduced safety.
- [ ] **Step 2: Ensure the key is included in config loading/saving paths**
Add `allow_local_remote_addr` anywhere the page enumerates config keys or groups persisted config values.
- [ ] **Step 3: Run frontend build**
Run: `pnpm run build`
Expected: PASS
- [ ] **Step 4: Run frontend lint**
Run: `pnpm run lint`
Expected: 0 errors; existing warnings may remain.
### Task 4: Final Verification
**Files:**
- Verify only
- [ ] **Step 1: Re-run backend contracts for the toggle**
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
Expected: PASS
- [ ] **Step 2: Re-run full backend tests**
Run: `go test ./...`
Expected: PASS
- [ ] **Step 3: Re-run frontend build/lint**
Run: `pnpm run build && pnpm run lint`
Expected: Build passes, lint has no errors.
- [ ] **Step 4: Commit**
```bash
git add go-backend/internal/http/handler/mutations.go go-backend/tests/contract/forward_contract_test.go vite-frontend/src/pages/config.tsx docs/superpowers/specs/2026-04-26-allow-local-remote-addr-design.md docs/superpowers/plans/2026-04-26-allow-local-remote-addr.md
git commit -m "feat: add allow-local-remote-address toggle"
```
@@ -0,0 +1,849 @@
# flow/upload Batch Optimization Implementation Plan
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
**Goal:** Reduce `POST /flow/upload` database pressure by converting the hot path from per-item queries and per-item transactions to per-request aggregation, batched metadata reads, and batched writes, while preserving immediate quota disable / forward pause behavior inside the same upload.
**Architecture:** Parse one upload into a batch object in the handler layer, fetch one shared `forward+tunnel` metadata map, then reuse that map for flow accounting and tunnel metric aggregation. Replace `AddFlow` and `AddUserQuotaUsage` per-item transactions with one batched flow transaction and one batched quota transaction; run policy enforcement, orphan cleanup, and peer-share flow handling once per affected target instead of once per item.
**Tech Stack:** Go, net/http, GORM, SQLite/PostgreSQL, existing backend contract tests.
---
## File Map
- Create: `go-backend/internal/http/handler/flow_upload_batch.go`
Responsibility: request-scoped parsing, aggregation, and application of one `/flow/upload` batch.
- Create: `go-backend/internal/http/handler/flow_upload_batch_test.go`
Responsibility: unit coverage for batch aggregation semantics.
- Create: `go-backend/internal/store/repo/repository_flow_batch_test.go`
Responsibility: unit coverage for batched flow and quota persistence.
- Create: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
Responsibility: contract coverage that repeated items still accumulate correctly and still disable quota immediately.
- Modify: `go-backend/internal/http/handler/handler.go`
Responsibility: switch `/flow/upload` entrypoint to the new batch pipeline.
- Modify: `go-backend/internal/http/handler/tunnel_metrics_ingestion.go`
Responsibility: accept pre-aggregated forward deltas plus shared forward metadata instead of reparsing the raw items.
- Modify: `go-backend/internal/store/repo/repository.go`
Responsibility: add batched flow persistence primitives near the existing flow update code.
- Modify: `go-backend/internal/store/repo/repository_flow.go`
Responsibility: add shared flow-upload metadata query helpers.
- Modify: `go-backend/internal/store/repo/repository_user_quota.go`
Responsibility: add batched quota usage persistence that still returns normalized quota views for immediate enforcement.
---
### Task 1: Add Failing Tests For Batched flow/upload Semantics
**Files:**
- Create: `go-backend/internal/http/handler/flow_upload_batch_test.go`
- Create: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
- [ ] **Step 1: Write the failing handler unit test**
Create `go-backend/internal/http/handler/flow_upload_batch_test.go` with a unit test that locks in the new aggregation contract.
```go
package handler
import (
"testing"
"go-backend/internal/store/repo"
)
func TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets(t *testing.T) {
h := &Handler{}
metas := map[int64]repo.FlowUploadForwardMeta{
20: {
ForwardID: 20,
TunnelID: 1,
TrafficRatio: 2,
TunnelFlow: 3,
},
}
batch := h.buildFlowUploadBatch([]flowItem{
{N: "20_2_10", U: 70, D: 50},
{N: "20_2_10_tcp", U: 40, D: 30},
{N: "99_2_10", U: 12, D: 8},
{N: "fed_svc_17", U: 9, D: 1},
}, metas)
if len(batch.flowDeltas) != 1 {
t.Fatalf("expected 1 flow delta, got %d", len(batch.flowDeltas))
}
delta := batch.flowDeltas[0]
if delta.ForwardID != 20 || delta.UserID != 2 || delta.UserTunnelID != 10 {
t.Fatalf("unexpected flow delta identity: %#v", delta)
}
if delta.InFlow != 480 || delta.OutFlow != 660 {
t.Fatalf("expected scaled flow in=480 out=660, got in=%d out=%d", delta.InFlow, delta.OutFlow)
}
if batch.quotaUsage[2] != 1140 {
t.Fatalf("expected quota usage 1140, got %d", batch.quotaUsage[2])
}
if len(batch.policyTargets) != 1 {
t.Fatalf("expected 1 policy target, got %d", len(batch.policyTargets))
}
if batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
t.Fatalf("unexpected policy target: %#v", batch.policyTargets[0])
}
traffic := batch.forwardTraffic[20]
if traffic.bytesIn != 80 || traffic.bytesOut != 110 {
t.Fatalf("expected raw traffic in=80 out=110, got in=%d out=%d", traffic.bytesIn, traffic.bytesOut)
}
if _, ok := batch.orphanServices["99_2_10"]; !ok {
t.Fatalf("expected orphan service cleanup target for 99_2_10")
}
if item, ok := batch.peerShareForwardItems["20_2_10"]; !ok || item.U != 110 || item.D != 80 {
t.Fatalf("expected merged peer-share forward item, got %#v ok=%v", item, ok)
}
if item, ok := batch.peerShareRuntimeItems[17]; !ok || item.U != 9 || item.D != 1 {
t.Fatalf("expected merged peer-share runtime item, got %#v ok=%v", item, ok)
}
}
```
- [ ] **Step 2: Run the handler unit test to verify RED**
Run:
```bash
go test ./internal/http/handler -run TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets -v
```
Expected: FAIL because `FlowUploadForwardMeta`, `buildFlowUploadBatch`, and the new batch fields do not exist yet.
- [ ] **Step 3: Write the contract test that guards current behavior**
Create `go-backend/tests/contract/flow_upload_batch_contract_test.go` so the optimization cannot weaken same-request quota enforcement.
```go
package contract_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"go-backend/internal/store/model"
)
func TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately(t *testing.T) {
secret := "monitoring-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now()
nowMs := now.UnixMilli()
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
monthKey := int64(now.Year()*100 + int(now.Month()))
const bytesPerGB = int64(1024 * 1024 * 1024)
node := &model.Node{Name: "node-1", Secret: "node-secret", ServerIP: "127.0.0.1", Port: "10000-10010", TCPListenAddr: "[::]", UDPListenAddr: "[::]", CreatedTime: nowMs, Status: 1}
if err := repo.DB().Create(node).Error; err != nil {
t.Fatalf("seed node: %v", err)
}
if err := repo.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'flow_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
tunnel := &model.Tunnel{Name: "tunnel-1", TrafficRatio: 1.0, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
if err := repo.DB().Create(tunnel).Error; err != nil {
t.Fatalf("seed tunnel: %v", err)
}
if err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnel.ID).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
forward := &model.Forward{ID: 20, UserID: 2, UserName: "flow_user", Name: "forward-20", TunnelID: tunnel.ID, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
if err := repo.DB().Create(forward).Error; err != nil {
t.Fatalf("seed forward: %v", err)
}
if err := repo.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 1, 0, ?, ?, ?, ?, 0, 0, '', ?, ?)`, bytesPerGB-100, bytesPerGB-100, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user_quota: %v", err)
}
body, err := json.Marshal([]map[string]interface{}{
{"n": "20_2_10", "u": 70, "d": 50},
{"n": "20_2_10_tcp", "u": 40, "d": 30},
})
if err != nil {
t.Fatalf("marshal body: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
if res.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", res.Code)
}
if got := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = 20`); got != 0 {
t.Fatalf("expected forward paused immediately, got status=%d", got)
}
if got := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`); got != 1 {
t.Fatalf("expected quota disabled flag=1, got %d", got)
}
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 80 {
t.Fatalf("expected forward in_flow=80, got %d", got)
}
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 110 {
t.Fatalf("expected forward out_flow=110, got %d", got)
}
metrics, err := repo.GetTunnelMetrics(tunnel.ID, 0, nowMs+60_000)
if err != nil {
t.Fatalf("get tunnel metrics: %v", err)
}
if len(metrics) != 1 || metrics[0].BytesIn != 80 || metrics[0].BytesOut != 110 {
t.Fatalf("expected one aggregated metric row, got %#v", metrics)
}
}
```
- [ ] **Step 4: Run the contract test to verify the same-request guard stays green or reveals an existing regression**
Run:
```bash
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
```
Expected: this test may already PASS before the refactor because it locks in existing external behavior. Keep it either way; it is the guardrail for the optimization.
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
```bash
git add go-backend/internal/http/handler/flow_upload_batch_test.go go-backend/tests/contract/flow_upload_batch_contract_test.go
git commit -m "test: cover flow upload batch semantics"
```
---
### Task 2: Add Batched Repository Primitives
**Files:**
- Modify: `go-backend/internal/store/repo/repository_flow.go`
- Modify: `go-backend/internal/store/repo/repository.go`
- Modify: `go-backend/internal/store/repo/repository_user_quota.go`
- Create: `go-backend/internal/store/repo/repository_flow_batch_test.go`
- [ ] **Step 1: Write the failing repository tests**
Create `go-backend/internal/store/repo/repository_flow_batch_test.go` with coverage for both the shared metadata query and the batched counter/quota writes.
```go
package repo
import (
"path/filepath"
"testing"
"time"
)
func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "flow-batch.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := r.DB().Exec(`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
metas, err := r.GetFlowUploadForwardMetas([]int64{20, 99})
if err != nil {
t.Fatalf("get metas: %v", err)
}
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
t.Fatalf("unexpected meta for forward 20: %#v", metas[20])
}
if _, ok := metas[99]; ok {
t.Fatalf("did not expect meta for missing forward 99")
}
err = r.ApplyFlowUploadDeltasBatch([]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}})
if err != nil {
t.Fatalf("apply flow batch: %v", err)
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
t.Fatalf("expected forward in_flow=480, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
t.Fatalf("expected user out_flow=660, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
}
}
func TestAddUserQuotaUsageBatchReturnsNormalizedViews(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "quota-batch.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
views, err := r.AddUserQuotaUsageBatch(map[int64]int64{2: 1140}, now)
if err != nil {
t.Fatalf("batch quota update: %v", err)
}
if views[2] == nil || views[2].DailyUsedBytes != 1140 || views[2].MonthlyUsedBytes != 1140 {
t.Fatalf("unexpected quota view: %#v", views[2])
}
}
func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...interface{}) int64 {
t.Helper()
var value int64
if err := r.DB().Raw(query, args...).Row().Scan(&value); err != nil {
t.Fatalf("query %q failed: %v", query, err)
}
return value
}
```
- [ ] **Step 2: Run the repository tests to verify RED**
Run:
```bash
go test ./internal/store/repo -run 'TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch|TestAddUserQuotaUsageBatchReturnsNormalizedViews' -v
```
Expected: FAIL because `GetFlowUploadForwardMetas`, `ApplyFlowUploadDeltasBatch`, `FlowUploadCounterDelta`, and `AddUserQuotaUsageBatch` do not exist yet.
- [ ] **Step 3: Implement shared flow-upload metadata and batched persistence**
Update `go-backend/internal/store/repo/repository_flow.go`, `repository.go`, and `repository_user_quota.go` with the following concrete APIs. Add `sort` to the `repository_user_quota.go` import list.
```go
// repository_flow.go
type FlowUploadForwardMeta struct {
ForwardID int64
TunnelID int64
TrafficRatio float64
TunnelFlow int64
}
func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]FlowUploadForwardMeta, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if len(forwardIDs) == 0 {
return map[int64]FlowUploadForwardMeta{}, nil
}
ids := make([]int64, 0, len(forwardIDs))
seen := make(map[int64]struct{}, len(forwardIDs))
for _, id := range forwardIDs {
if id <= 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
ids = append(ids, id)
}
type row struct {
ForwardID int64 `gorm:"column:forward_id"`
TunnelID int64 `gorm:"column:tunnel_id"`
TrafficRatio float64 `gorm:"column:traffic_ratio"`
TunnelFlow int64 `gorm:"column:tunnel_flow"`
}
var rows []row
err := r.db.Table("forward AS f").
Select("f.id AS forward_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
Joins("JOIN tunnel t ON t.id = f.tunnel_id").
Where("f.id IN ?", ids).
Scan(&rows).Error
if err != nil {
return nil, err
}
out := make(map[int64]FlowUploadForwardMeta, len(rows))
for _, row := range rows {
if row.TunnelFlow <= 0 {
row.TunnelFlow = 1
}
if row.TrafficRatio <= 0 {
row.TrafficRatio = 1
}
out[row.ForwardID] = FlowUploadForwardMeta{ForwardID: row.ForwardID, TunnelID: row.TunnelID, TrafficRatio: row.TrafficRatio, TunnelFlow: row.TunnelFlow}
}
return out, nil
}
```
```go
// repository.go
type FlowUploadCounterDelta struct {
ForwardID int64
UserID int64
UserTunnelID int64
InFlow int64
OutFlow int64
}
func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if len(deltas) == 0 {
return nil
}
forwardTotals := make(map[int64][2]int64, len(deltas))
userTotals := make(map[int64][2]int64, len(deltas))
userTunnelTotals := make(map[int64][2]int64, len(deltas))
for _, delta := range deltas {
if delta.ForwardID > 0 {
current := forwardTotals[delta.ForwardID]
current[0] += delta.InFlow
current[1] += delta.OutFlow
forwardTotals[delta.ForwardID] = current
}
if delta.UserID > 0 {
current := userTotals[delta.UserID]
current[0] += delta.InFlow
current[1] += delta.OutFlow
userTotals[delta.UserID] = current
}
if delta.UserTunnelID > 0 {
current := userTunnelTotals[delta.UserTunnelID]
current[0] += delta.InFlow
current[1] += delta.OutFlow
userTunnelTotals[delta.UserTunnelID] = current
}
}
return r.db.Transaction(func(tx *gorm.DB) error {
for forwardID, total := range forwardTotals {
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1])}).Error; err != nil {
return err
}
}
for userID, total := range userTotals {
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1])}).Error; err != nil {
return err
}
}
for userTunnelID, total := range userTunnelTotals {
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1])}).Error; err != nil {
return err
}
}
return nil
})
}
```
```go
// repository_user_quota.go
func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if len(usages) == 0 {
return map[int64]*model.UserQuotaView{}, nil
}
result := make(map[int64]*model.UserQuotaView, len(usages))
err := r.db.Transaction(func(tx *gorm.DB) error {
userIDs := make([]int64, 0, len(usages))
for userID := range usages {
if userID > 0 {
userIDs = append(userIDs, userID)
}
}
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
for _, userID := range userIDs {
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
if err != nil {
return err
}
applyUserQuotaWindowRoll(q, now)
if usages[userID] > 0 {
q.DailyUsedBytes += usages[userID]
q.MonthlyUsedBytes += usages[userID]
}
q.UpdatedTime = now.UnixMilli()
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{"daily_used_bytes": q.DailyUsedBytes, "monthly_used_bytes": q.MonthlyUsedBytes, "day_key": q.DayKey, "month_key": q.MonthKey, "updated_time": q.UpdatedTime}).Error; err != nil {
return err
}
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
}
return nil
})
if err != nil {
return nil, err
}
return result, nil
}
```
- [ ] **Step 4: Run the repository tests to verify GREEN**
Run:
```bash
go test ./internal/store/repo -run 'TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch|TestAddUserQuotaUsageBatchReturnsNormalizedViews' -v
```
Expected: PASS.
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
```bash
git add go-backend/internal/store/repo/repository.go go-backend/internal/store/repo/repository_flow.go go-backend/internal/store/repo/repository_user_quota.go go-backend/internal/store/repo/repository_flow_batch_test.go
git commit -m "refactor: batch flow upload persistence"
```
---
### Task 3: Refactor flow/upload To Use One Parsed Batch
**Files:**
- Create: `go-backend/internal/http/handler/flow_upload_batch.go`
- Modify: `go-backend/internal/http/handler/handler.go`
- Modify: `go-backend/internal/http/handler/tunnel_metrics_ingestion.go`
- Modify: `go-backend/internal/http/handler/flow_upload_batch_test.go`
- Modify: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
- [ ] **Step 1: Write the new handler batch implementation**
Create `go-backend/internal/http/handler/flow_upload_batch.go` and move the request-scoped aggregation there.
```go
package handler
import (
"log"
"sort"
"strings"
"time"
"go-backend/internal/store/repo"
)
type flowPolicyTarget struct {
UserID int64
UserTunnelID int64
}
type flowUploadBatch struct {
flowDeltas []repo.FlowUploadCounterDelta
quotaUsage map[int64]int64
policyTargets []flowPolicyTarget
forwardTraffic map[int64]tunnelTrafficDelta
orphanServices map[string]struct{}
peerShareForwardItems map[string]flowItem
peerShareRuntimeItems map[int64]flowItem
}
func (h *Handler) buildFlowUploadBatch(items []flowItem, metas map[int64]repo.FlowUploadForwardMeta) flowUploadBatch {
batch := flowUploadBatch{
quotaUsage: make(map[int64]int64),
forwardTraffic: make(map[int64]tunnelTrafficDelta),
orphanServices: make(map[string]struct{}),
peerShareForwardItems: make(map[string]flowItem),
peerShareRuntimeItems: make(map[int64]flowItem),
}
policySeen := map[flowPolicyTarget]struct{}{}
flowSeen := map[int64]int{}
for _, item := range items {
serviceName := strings.TrimSpace(item.N)
if serviceName == "" || serviceName == "web_api" {
continue
}
if runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName); ok {
merged := batch.peerShareRuntimeItems[runtimeID]
merged.N = serviceName
merged.U += item.U
merged.D += item.D
batch.peerShareRuntimeItems[runtimeID] = merged
continue
}
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
if !ok {
continue
}
meta, exists := metas[forwardID]
if !exists {
batch.orphanServices[serviceName] = struct{}{}
continue
}
raw := batch.forwardTraffic[forwardID]
raw.bytesIn += item.D
raw.bytesOut += item.U
batch.forwardTraffic[forwardID] = raw
scaledIn := int64(float64(item.D)*meta.TrafficRatio) * meta.TunnelFlow
scaledOut := int64(float64(item.U)*meta.TrafficRatio) * meta.TunnelFlow
if idx, ok := flowSeen[forwardID]; ok {
batch.flowDeltas[idx].InFlow += scaledIn
batch.flowDeltas[idx].OutFlow += scaledOut
} else {
flowSeen[forwardID] = len(batch.flowDeltas)
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{ForwardID: forwardID, UserID: userID, UserTunnelID: userTunnelID, InFlow: scaledIn, OutFlow: scaledOut})
}
batch.quotaUsage[userID] += scaledIn + scaledOut
target := flowPolicyTarget{UserID: userID, UserTunnelID: userTunnelID}
if _, seen := policySeen[target]; !seen {
policySeen[target] = struct{}{}
batch.policyTargets = append(batch.policyTargets, target)
}
merged := batch.peerShareForwardItems[normalizeForwardRuntimeServiceName(serviceName)]
merged.N = normalizeForwardRuntimeServiceName(serviceName)
merged.U += item.U
merged.D += item.D
batch.peerShareForwardItems[normalizeForwardRuntimeServiceName(serviceName)] = merged
}
sort.Slice(batch.policyTargets, func(i, j int) bool {
if batch.policyTargets[i].UserID == batch.policyTargets[j].UserID {
return batch.policyTargets[i].UserTunnelID < batch.policyTargets[j].UserTunnelID
}
return batch.policyTargets[i].UserID < batch.policyTargets[j].UserID
})
return batch
}
func (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now time.Time) {
if h == nil || h.repo == nil {
return
}
if err := h.repo.ApplyFlowUploadDeltasBatch(batch.flowDeltas); err != nil {
log.Printf("flow upload write failed op=flow.batch_apply node_id=%d err=%v", nodeID, err)
return
}
quotaViews, err := h.repo.AddUserQuotaUsageBatch(batch.quotaUsage, now)
if err != nil {
log.Printf("flow upload write failed op=quota.batch_apply node_id=%d err=%v", nodeID, err)
return
}
for userID, quota := range quotaViews {
h.enforceUserQuotaIfNeeded(userID, quota)
}
for _, target := range batch.policyTargets {
if target.UserID <= 0 || target.UserTunnelID <= 0 {
continue
}
h.enforceFlowPolicies(target.UserID, target.UserTunnelID)
}
for serviceName := range batch.orphanServices {
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
}
for serviceName, item := range batch.peerShareForwardItems {
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
if ok {
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
}
}
for runtimeID, item := range batch.peerShareRuntimeItems {
h.processPeerShareFlow(runtimeID, item)
}
}
```
- [ ] **Step 2: Switch the `/flow/upload` entrypoint and tunnel metric ingestion to the shared batch**
Modify `handler.go` and `tunnel_metrics_ingestion.go` so the raw JSON is parsed once and the same forward metadata powers both flow counters and tunnel metrics.
```go
// handler.go
func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
secret := r.URL.Query().Get("secret")
node, _ := h.repo.GetNodeBySecret(secret)
if node == nil {
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
_, _ = w.Write([]byte("ok"))
return
}
raw, err := readAndDecryptFlowBody(r.Body, secret)
if err == nil && strings.TrimSpace(raw) != "" {
var items []flowItem
if json.Unmarshal([]byte(raw), &items) == nil {
now := time.Now()
forwardIDs := collectFlowUploadForwardIDs(items)
metas, metaErr := h.repo.GetFlowUploadForwardMetas(forwardIDs)
if metaErr != nil {
log.Printf("flow upload metadata lookup failed node_id=%d err=%v", node.ID, metaErr)
metas = map[int64]repo.FlowUploadForwardMeta{}
}
batch := h.buildFlowUploadBatch(items, metas)
h.recordTunnelMetricsFromForwardBatch(node.ID, batch.forwardTraffic, metas, now.UnixMilli())
h.applyFlowUploadBatch(node.ID, batch, now)
}
}
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
_, _ = w.Write([]byte("ok"))
}
```
```go
// tunnel_metrics_ingestion.go
func collectFlowUploadForwardIDs(items []flowItem) []int64 {
ids := make([]int64, 0, len(items))
seen := make(map[int64]struct{}, len(items))
for _, item := range items {
forwardID, _, _, ok := parseFlowServiceIDs(strings.TrimSpace(item.N))
if !ok || forwardID <= 0 {
continue
}
if _, exists := seen[forwardID]; exists {
continue
}
seen[forwardID] = struct{}{}
ids = append(ids, forwardID)
}
return ids
}
func (h *Handler) recordTunnelMetricsFromForwardBatch(nodeID int64, forwardDeltas map[int64]tunnelTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta, nowMs int64) {
if h == nil || h.repo == nil || nodeID <= 0 || len(forwardDeltas) == 0 {
return
}
bucketTs := unixMilliBucketMinute(nowMs)
if bucketTs <= 0 {
return
}
tunnelAgg := make(map[int64]tunnelTrafficDelta)
for forwardID, delta := range forwardDeltas {
meta, ok := metas[forwardID]
if !ok || meta.TunnelID <= 0 {
continue
}
current := tunnelAgg[meta.TunnelID]
current.bytesIn += delta.bytesIn
current.bytesOut += delta.bytesOut
tunnelAgg[meta.TunnelID] = current
}
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
for tunnelID, delta := range tunnelAgg {
if delta.bytesIn == 0 && delta.bytesOut == 0 {
continue
}
metrics = append(metrics, &model.TunnelMetric{TunnelID: tunnelID, NodeID: nodeID, Timestamp: bucketTs, BytesIn: delta.bytesIn, BytesOut: delta.bytesOut})
}
if len(metrics) == 0 {
return
}
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
return
}
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
}
```
- [ ] **Step 3: Run focused handler and contract tests to verify GREEN**
Run:
```bash
go test ./internal/http/handler -run TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets -v
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
```
Expected: PASS.
- [ ] **Step 4: Run the full backend suite**
Run:
```bash
go test ./...
```
Expected: PASS across the backend module.
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
```bash
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/tunnel_metrics_ingestion.go go-backend/internal/http/handler/flow_upload_batch.go go-backend/internal/http/handler/flow_upload_batch_test.go go-backend/tests/contract/flow_upload_batch_contract_test.go
git commit -m "refactor: batch flow upload processing"
```
---
### Task 4: Final Verification And Performance Sanity Check
**Files:**
- Modify: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
- [ ] **Step 1: Add a same-batch duplicate-item stress assertion**
Extend the contract test with a second request that repeats the same service name multiple times and assert the counters advance by exactly the summed amount.
```go
body, err = json.Marshal([]map[string]interface{}{
{"n": "20_2_10", "u": 10, "d": 20},
{"n": "20_2_10", "u": 10, "d": 20},
{"n": "20_2_10_tcp", "u": 10, "d": 20},
})
if err != nil {
t.Fatalf("marshal body: %v", err)
}
req = httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
res = httptest.NewRecorder()
router.ServeHTTP(res, req)
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 140 {
t.Fatalf("expected forward in_flow=140 after second request, got %d", got)
}
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 140 {
t.Fatalf("expected forward out_flow=140 after second request, got %d", got)
}
```
- [ ] **Step 2: Run the targeted contract test again**
Run:
```bash
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
```
Expected: PASS.
- [ ] **Step 3: Re-run the full backend suite before claiming completion**
Run:
```bash
go test ./...
```
Expected: PASS.
- [ ] **Step 4: Optional local profiling sanity check**
Run a short local comparison before and after the change with the same repeated flow payload.
```bash
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -count=10
```
Expected: the test remains stable across repeated runs and does not introduce flakiness.
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
```bash
git add go-backend/tests/contract/flow_upload_batch_contract_test.go
git commit -m "test: harden flow upload batch regression coverage"
```
@@ -0,0 +1,133 @@
# 允许转发到本地地址开关设计
**日期**: 2026-04-26
**状态**: 待审核
**作者**: AI Assistant
## 概述
新增一个全局设置开关,控制规则目标地址是否允许指向本地/内网地址。默认关闭,保持当前安全策略不变;开启后,规则创建和编辑时允许将目标地址设置为 `127.0.0.1`、`10.x.x.x`、`172.16-31.x.x`、`192.168.x.x` 等本地或私网地址。
## 背景
当前后端在规则创建和编辑时会调用 `IsSafeRemoteAddr()`,统一禁止目标地址指向本地/内网地址,用来降低 SSRF / 开放代理风险。这一行为是全局硬编码的,无法按部署场景调整。
有些用户需要把规则转发到本机或内网服务,因此需要一个显式、全局的开关来放宽这条限制。
## 目标
1. 在设置页提供一个全局开关控制该行为。
2. 默认关闭,不改变现有安全默认值。
3. 开启后,规则创建和编辑允许本地/内网目标地址。
4. 不影响其他安全校验和其他业务流程。
## 影响范围
### 后端
- `go-backend/internal/http/handler/security_utils.go`
- `go-backend/internal/http/handler/mutations.go`
- `go-backend/internal/http/handler/handler.go`
### 前端
- `vite-frontend/src/pages/config.tsx`
### 测试
- `go-backend/tests/contract/forward_contract_test.go` 或新增独立 contract test
## 详细设计
### 1. 配置存储
使用现有 `vite_config` 表新增一个配置项:
| name | value | 说明 |
|------|-------|------|
| `allow_local_remote_addr` | `"1"` / `"0"` | 是否允许规则目标地址指向本地/内网地址 |
约定:
- 未配置时按 `"0"` 处理
- `"1"` 表示允许
- 其他值一律按关闭处理
### 2. 后端行为
新增一个轻量辅助函数,用于读取该配置开关:
```go
func (h *Handler) allowLocalRemoteAddr() bool {
if h == nil || h.repo == nil {
return false
}
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
if err != nil || cfg == nil {
return false
}
return strings.TrimSpace(cfg.Value) == "1"
}
```
在以下路径中应用:
- `forwardCreate`
- `forwardUpdate`
行为改为:
- 当开关关闭时,继续执行 `IsSafeRemoteAddr(remoteAddr)`
- 当开关开启时,跳过这条“本地/内网地址禁止”校验
这样可以把改动范围限定在规则创建/编辑,不改变其他依赖 `IsSafeRemoteAddr()` 的场景。
### 3. 前端设置页
在 `vite-frontend/src/pages/config.tsx` 增加一个全局开关配置项。
建议文案:
- 标签:`允许转发到本地地址`
- 描述:`开启后,规则目标地址可指向 127.0.0.1、10.x.x.x、172.16-31.x.x、192.168.x.x 等本地或内网地址。默认关闭以降低开放代理风险。`
控件类型:
- 使用现有设置页的布尔开关模式
默认显示策略:
- 不依赖其他配置项
- 直接显示在设置页的网络/安全相关区域;若现有页面没有单独分区,则先按现有配置项组织方式加入即可
### 4. 错误与兼容性
关闭开关时:
- 保持现有错误行为,继续阻止本地/内网地址
开启开关时:
- 仅放开“本地/内网地址禁止”这条限制
- 仍保留地址格式解析失败等其他错误
### 5. 测试
需要补两类后端契约测试:
1. 开关关闭时拒绝本地/内网地址
- 创建规则时使用本地/内网地址
- 断言接口返回非 0 code
2. 开关开启时允许本地/内网地址
- 先写入 `vite_config(name=allow_local_remote_addr, value=1)`
- 创建或更新规则时使用相同地址
- 断言接口成功
建议至少覆盖:
- create 路径
- update 路径
- 多目标地址输入(逗号或换行分隔)中包含本地地址时的行为
## 风险与约束
1. 该开关会降低默认安全防护,应明确标注风险。
2. 这是全局开关,不做用户级或规则级细分控制。
3. 该开关只影响规则目标地址校验,不影响其他独立的安全策略。
## 推荐实施顺序
1. 先补失败的后端契约测试
2. 实现后端配置读取与创建/更新分支控制
3. 在设置页增加开关
4. 跑后端测试与前端构建验证
@@ -0,0 +1,321 @@
# 规则每 IP 连接数与限速设计
**日期**: 2026-04-27
**状态**: 待审核
**作者**: AI Assistant
## 概述
在转发规则的高级设置中新增两类每客户端 IP 限制:每 IP 最大连接数、每 IP 带宽限速。保留现有总量限制语义不变,新增字段只在用户显式配置时生效。
实现优先复用 GOST 已有能力:`climiters` 的 `$$ N` 表示每个客户端 IP 独立最大连接数;`limiters` 支持 IP/CIDR 级带宽桶,可用 `0.0.0.0/0` 和 `::/0` 实现默认覆盖所有 IPv4/IPv6 客户端的每 IP 带宽限速。
## 背景
当前 FLVX 已经支持规则级最大连接数和规则级限速,但这两个限制都是规则总量:
- `maxConn` 下发为 GOST `climiters` 的 `$ N`,限制整条规则的总并发连接数。
- `speedId` 下发为 GOST `limiters` 的 `$ in out`,限制整条规则的总带宽。
用户需要的是按客户端 IP 隔离的限制,例如每个 IP 最多 5 个连接、每个 IP 最多 10 Mbps,而不是所有客户端共享同一个总量。
## GOST 能力确认
### 连接数限制
`go-gost/x/limiter/conn/conn.go` 已内置以下语义:
| Key | 含义 |
|-----|------|
| `$` | 全局连接数限制,所有客户端共享一个 limiter |
| `$$` | 每个客户端 IP 独立连接数限制,每个 IP 创建自己的 limiter |
| `IP` / `CIDR` | 指定 IP 或 CIDR 的连接数限制 |
因此每 IP 连接数无需新增 agent 限制器,只需后端下发 `$$ N`。
### 带宽限制
`go-gost/x/limiter/traffic/traffic.go` 已内置以下语义:
| Key | 含义 |
|-----|------|
| `$` | 服务级总带宽限制 |
| `$$` | 连接级带宽限制 |
| `IP` / `CIDR` | 客户端 IP 或 CIDR 级带宽限制 |
CIDR 级限制使用 generator,为命中的客户端 IP 创建独立 limiter。使用 `0.0.0.0/0` 和 `::/0` 可以覆盖所有 IPv4/IPv6 客户端,实现每 IP 带宽限速。
### 现有缺口
TCP listener 已在 Accept 后用客户端地址包装连接级 traffic limiter,路径可用于每 IP 带宽。UDP listener 当前只在 PacketConn 上应用服务级 limiter,没有在 `Accept()` 后按客户端 UDP pseudo-connection 包装 limiter,也没有挂接 connection limiter。因此要让 UDP 与 TCP 语义一致,需要补齐 UDP listener 的 per-client wrapper。
## 目标
1. 保留现有 `maxConn` 和 `speedId` 的总量语义。
2. 在规则上新增每 IP 最大连接数。
3. 在规则上新增每 IP 带宽限速。
4. 同一规则允许同时配置总量限制和每 IP 限制。
5. 普通用户不能设置或修改限速规则字段,保持现有权限模型。
6. TCP 和 UDP 入口都尽量遵循相同限制语义。
## 非目标
1. 不新增按用户组、节点组、国家地区、ASN 的限制。
2. 不新增请求频率限制;本次“每个 IP 限速”指带宽限速,不是新建连接频率。
3. 不改变已有 speed limit 规则表的单位和含义。
4. 不把用户级默认最大连接数改成每 IP 语义;用户级 `maxConn` 继续作为默认总连接数。
## 数据模型
在 `forward` 表新增两个字段:
| 字段 | 类型 | 默认 | 说明 |
|------|------|------|------|
| `ip_max_conn` | int | `0` | 每 IP 最大连接数,`0` 表示不启用 |
| `ip_speed_id` | nullable int64 | `NULL` | 每 IP 带宽限速规则 ID,`NULL` 表示不启用 |
Go 模型新增:
```go
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
```
字段会通过现有 auto-migrate 机制创建,保持 SQLite/PostgreSQL 兼容,不使用 SQLite 不兼容的 GORM tags。
## API 行为
### 创建规则
`/forward/create` 新增入参:
```json
{
"ipMaxConn": 5,
"ipSpeedId": 123
}
```
规则:
- `ipMaxConn` 缺省或小于等于 `0` 时按 `0` 存储,不启用每 IP 连接数限制。
- `ipSpeedId` 缺省或不存在时存为 `NULL`,不启用每 IP 带宽限速。
- `ipSpeedId` 指向不存在的限速规则时按 `NULL` 处理,沿用现有 `speedId` 的容错策略。
- 普通用户提交非空 `ipSpeedId` 时返回错误,保持与 `speedId` 一致的权限边界。
### 更新规则
`/forward/update` 新增入参:
```json
{
"ipMaxConn": 5,
"ipSpeedId": 123
}
```
规则:
- 未提交 `ipMaxConn` 时保留原值;提交空值或 `0` 时清除每 IP 连接数限制。
- 未提交 `ipSpeedId` 时保留原值;提交 `null` 时清除每 IP 带宽限速。
- 普通用户不能把 `ipSpeedId` 改成不同的非空值。
- 更新后重新同步运行时服务和 limiter。
### 列表返回
`/forward/list` 返回项新增:
```json
{
"ipMaxConn": 5,
"ipSpeedId": 123,
"ipSpeedLimitName": "每IP 10Mbps"
}
```
`ipSpeedLimitName` 可选,但建议返回,便于前端显示缺失或已删除的限速规则。
## 后端运行时同步
### 连接数限制器
将现有连接限制器构建从单一总量扩展为组合规则。
当前行为:
```json
{
"name": "rule_conn_limit_42",
"limits": ["$ 100"]
}
```
新增行为:
```json
{
"name": "rule_conn_limit_42",
"limits": ["$ 100", "$$ 5"]
}
```
规则:
- `maxConn > 0` 时追加 `$ maxConn`。
- `ipMaxConn > 0` 时追加 `$$ ipMaxConn`。
- 如果规则未配置 `maxConn` 且用户有 `MaxConn > 0`,继续继承用户级总连接数,追加 `$ user.MaxConn`。
- 如果两者都没有,则不下发 `climiter`,服务不引用 `climiter`。
- limiter 名称继续优先使用 `rule_conn_limit_<forwardID>`;只有用户级默认总连接数且规则没有任何连接限制时可继续使用 `user_conn_limit_<userID>`,避免不必要的 per-rule limiter。
### 带宽限制器
将现有规则限速从单一 `speedId` 扩展为组合 limiter。
当前行为:
```json
{
"name": "123",
"limits": ["$ 1.3MB 1.3MB"]
}
```
新增每 IP 行为:
```json
{
"name": "rule_traffic_limit_42",
"limits": [
"$ 1.3MB 1.3MB",
"0.0.0.0/0 1.3MB 1.3MB",
"::/0 1.3MB 1.3MB"
]
}
```
规则:
- 只有总量 `speedId` 时,保持现有名称和下发路径,服务继续引用 `speedId` 字符串。
- 只有每 IP `ipSpeedId` 时,创建 `rule_traffic_limit_<forwardID>`,只包含 IPv4/IPv6 CIDR 行。
- 总量和每 IP 同时存在时,创建 `rule_traffic_limit_<forwardID>`,同时包含 `$` 和 CIDR 行。
- 如果规则没有 `speedId`,则总量仍可继承 user tunnel 的 `speedId`,保持现有 fallback 语义;当继承的总量限速与 `ipSpeedId` 同时存在时,也使用 `rule_traffic_limit_<forwardID>` 组合 limiter。
- 每 IP 限速不从 user tunnel 继承,只由规则字段控制。
- `AddLimiters` 失败且提示已存在时,使用 `UpdateLimiters` 更新。
### 服务配置
`buildForwardServiceConfigs` 需要从当前 `limiterID *int64` / `cLimiterName string` 扩展为更明确的运行时限制描述,例如:
```go
type forwardRuntimeLimiters struct {
TrafficLimiter string
ConnLimiter string
}
```
服务配置只关心最终引用的 limiter 名称:
- `service["limiter"] = runtimeLimiters.TrafficLimiter`
- `service["climiter"] = runtimeLimiters.ConnLimiter`
这样可以把“如何构建 limiter payload”的逻辑和“如何构建 service JSON”的逻辑分开。
## Agent/GOST 调整
### WebSocket 命令
当前 agent WebSocket 已支持:
- `AddLimiters` / `UpdateLimiters` / `DeleteLimiters`
- `AddCLimiters` / `UpdateCLimiters` / `DeleteCLimiters`
本设计无需新增命令类型。
### UDP listener
补齐 `go-gost/x/listener/udp/listener.go` 的 `Accept()` 包装逻辑,使 UDP pseudo-connection 与 TCP listener 一致:
- 对 `l.options.ConnLimiter` 按客户端地址应用连接数限制。
- 对 `l.options.TrafficLimiter` 按 `conn.RemoteAddr().String()` 应用连接级 traffic wrapper。
需要注意 UDP pseudo-connection 的生命周期由内部 UDP listener 的 TTL/keepalive 控制;connection limiter 必须在 pseudo-connection 关闭时释放计数。
## 前端设计
在 `vite-frontend/src/pages/forward.tsx` 的规则高级设置中新增两个控件:
1. `每 IP 最大连接数`
- 类型:number input。
- 文案:`每个客户端 IP 可同时建立的最大连接数;0 或空表示不限制。`
- 字段:`ipMaxConn`。
2. `每 IP 限速`
- 类型:Select,复用现有限速规则列表。
- 文案:`每个客户端 IP 独享该带宽限制;不选择表示不限制。`
- 字段:`ipSpeedId`。
- 只对管理员显示,保持与 `规则限速` 一致。
前端类型需要同步更新:
- `ForwardApiItem`
- `ForwardMutationPayload`
- `ForwardForm` 或页面内等价类型
## 错误处理与兼容性
1. 旧数据默认 `ip_max_conn=0`、`ip_speed_id=NULL`,行为与当前版本一致。
2. 现有 agent 已支持 limiter 命令和 GOST limiter 语法;发布时需要包含 UDP 修复,才能让 TCP/UDP 都获得完整语义。
3. 节点离线时沿用现有 warning 行为,规则仍可保存,在线节点跳过下发。
4. 如果每 IP speed limit ID 被删除,更新时按 `NULL` 处理,列表页可提示或自动清除,和现有 `speedId` 行为一致。
5. 如果 IPv6 CIDR 在某些监听路径未命中,IPv4 行仍正常生效;测试应覆盖 IPv4,IPv6 通过 payload 合同保证下发。
## 测试计划
### 后端 contract 测试
新增或扩展 `go-backend/tests/contract/max_conn_limit_contract_test.go`:
1. 创建规则时设置 `ipMaxConn=5`,断言 `AddCLimiters` payload 包含 `$$ 5`。
2. 同时设置 `maxConn=100` 和 `ipMaxConn=5`,断言 payload 包含 `$ 100` 和 `$$ 5`。
3. 用户级 `MaxConn` 存在且规则 `ipMaxConn=5` 时,断言 payload 包含 `$ userMaxConn` 和 `$$ 5`。
新增每 IP 限速 contract 测试:
1. 创建规则时设置 `ipSpeedId`,断言 `AddLimiters` payload 包含 `0.0.0.0/0 ...` 和 `::/0 ...`。
2. 同时设置 `speedId` 和 `ipSpeedId`,断言组合 limiter 包含 `$ ...` 与两个 CIDR 行,服务引用 `rule_traffic_limit_<forwardID>`。
3. 普通用户提交 `ipSpeedId` 返回错误。
### Repository/API 测试
1. `CreateForwardTx`、`UpdateForward`、列表查询读写 `ip_max_conn` 和 `ip_speed_id`。
2. `/forward/list` 返回 `ipMaxConn`、`ipSpeedId`。
### GOST/x 测试
1. `go-gost/x/limiter/conn`:验证 `$$ N` 为不同 IP 创建独立 limiter。
2. `go-gost/x/limiter/traffic`:验证 `0.0.0.0/0` 为不同 IPv4 创建独立 limiter。
3. UDP listener:验证 Accept 返回的 UDP pseudo-connection 关闭后释放 connection limiter。
### 验证命令
```bash
(cd go-backend && go test ./...)
(cd go-gost/x && go test ./limiter/... ./listener/udp/...)
(cd vite-frontend && pnpm run build)
```
## 推荐实施顺序
1. 后端模型、repo DTO、API 字段读写。
2. 后端 limiter payload 构建与服务引用重构。
3. Contract 测试覆盖连接数和带宽 payload。
4. GOST UDP listener per-client wrapper 与相关测试。
5. 前端高级设置表单和类型更新。
6. 运行后端测试、GOST/x 相关测试、前端构建。
## 风险
1. UDP pseudo-connection 生命周期和 TCP 连接不同,连接数释放必须依赖 Close 包装正确执行。
2. 总带宽和每 IP 带宽组合时 limiter 名称从纯 speed ID 变为 rule-level 名称,需要确保更新已有规则时不会留下错误引用。
3. 旧节点如果没有 UDP wrapper 修复,TCP 生效但 UDP 每 IP 语义可能不完整;发布时应要求 agent 同步升级。
4. 每 IP 带宽是每个入口节点本地独立限制,不是跨节点全局聚合限制。
+1 -1
View File
@@ -90,7 +90,7 @@
| 表名 | Model | 特殊处理 |
|------|-------|----------|
| `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) |
| `forward` | `Forward` | |
| `forward` | `Forward` | 增加 `proxy_protocol` 字段 |
| `forward_port` | `ForwardPort` | |
| `node` | `Node` | |
| `speed_limit` | `SpeedLimit` | |
@@ -71,10 +71,11 @@ type RuntimeReleaseRoleRequest struct {
}
type RuntimeDiagnoseRequest struct {
IP string `json:"ip"`
Port int `json:"port"`
Count int `json:"count"`
Timeout int `json:"timeout"`
IP string `json:"ip"`
Port int `json:"port"`
Count int `json:"count"`
Timeout int `json:"timeout"`
Protocol string `json:"protocol"`
}
type RuntimeNodeCommandRequest struct {
+262 -26
View File
@@ -27,6 +27,16 @@ type nodeRecord = model.NodeRecord
type chainNodeRecord = model.ChainNodeRecord
type forwardRuntimeLimiters struct {
TrafficLimiter string
ConnLimiter string
}
type forwardLimiterConfig struct {
Name string
Limits []string
}
type diagnosisTarget struct {
Address string
IP string
@@ -42,6 +52,7 @@ type diagnosisWorkItem struct {
toNode chainNodeRecord
hasChainHop bool
ipPreference string
protocol string
}
type diagnosisExecOptions struct {
@@ -263,10 +274,31 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
speed = utSpeed
}
var ipSpeed *int
if forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0 {
if speedVal, err := h.repo.GetSpeedLimitSpeed(forward.IPSpeedID.Int64); err == nil && speedVal > 0 {
ipSpeed = &speedVal
}
}
serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
user, err := h.repo.GetUserByID(forward.UserID)
if err != nil {
return nil, err
}
userMaxConn := 0
if user != nil && user.MaxConn > 0 {
userMaxConn = user.MaxConn
}
connLimiterConfigs := buildConnLimiterConfigs(forward, userMaxConn)
for _, fp := range ports {
runtimeLimiters := forwardRuntimeLimiters{ConnLimiter: joinLimiterNames(connLimiterConfigs)}
trafficLimiterNames := make([]string, 0, 2)
if limiterID != nil && speed != nil {
totalLimiterName := strconv.FormatInt(*limiterID, 10)
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
// If the limiter push fails because the node is offline, skip it with a warning
if isNodeOfflineOrTimeoutError(err) {
@@ -280,13 +312,38 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
}
return nil, err
}
trafficLimiterNames = append(trafficLimiterNames, totalLimiterName)
}
if ipSpeed != nil {
ruleLimiterName := fmt.Sprintf("rule_traffic_limit_%d", forward.ID)
if err := h.ensureTrafficLimiterOnNode(fp.NodeID, ruleLimiterName, nil, ipSpeed); err != nil {
// If the limiter push fails because the node is offline, skip it with a warning
if isNodeOfflineOrTimeoutError(err) {
node, _ := h.getNodeRecord(fp.NodeID)
nodeName := fmt.Sprintf("%d", fp.NodeID)
if node != nil && strings.TrimSpace(node.Name) != "" {
nodeName = strings.TrimSpace(node.Name)
}
warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName))
continue
}
return nil, err
}
trafficLimiterNames = append(trafficLimiterNames, ruleLimiterName)
}
runtimeLimiters.TrafficLimiter = strings.Join(trafficLimiterNames, ",")
for _, connLimiterConfig := range connLimiterConfigs {
if err := h.ensureConnLimiterOnNode(fp.NodeID, connLimiterConfig); err != nil {
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
}
}
node, err := h.getNodeRecord(fp.NodeID)
if err != nil {
return nil, err
}
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID)
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), runtimeLimiters)
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
if err != nil && allowFallbackAdd && method == "UpdateService" {
if isNotFoundError(err) {
@@ -301,7 +358,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
}
if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isCannotAssignRequestedAddressError(err) {
var warning string
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID)
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, runtimeLimiters)
if err == nil && warning != "" {
warnings = append(warnings, warning)
}
@@ -327,7 +384,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
return warnings, nil
}
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64) (string, error) {
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, runtimeLimiters forwardRuntimeLimiters) (string, error) {
if h == nil || forward == nil || tunnel == nil || node == nil {
return "", errors.New("invalid bind fallback context")
}
@@ -344,7 +401,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne
}
time.Sleep(150 * time.Millisecond)
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID)
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", runtimeLimiters)
if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil {
return "", err
}
@@ -473,6 +530,9 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
nodeHandled, lastNotFoundErr, err := h.controlForwardServicesOnNode(fp.NodeID, bases, commandType)
if err != nil {
if isNodeOfflineOrTimeoutError(err) {
continue
}
return err
}
@@ -692,6 +752,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
}
ipPreference := h.repo.GetTunnelIPPreference(forward.TunnelID)
protocol := strings.ToLower(strings.TrimSpace(tunnel.Protocol))
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
workItems := make([]diagnosisWorkItem, 0, len(chainRows)*2+len(targets))
@@ -706,6 +767,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
targetIP: target.IP,
targetPort: target.Port,
description: description,
protocol: "tcp",
metadata: map[string]interface{}{
"fromChainType": 1,
},
@@ -723,6 +785,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(firstNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 1,
"toChainType": 2,
@@ -739,6 +802,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(outNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 1,
"toChainType": 3,
@@ -759,6 +823,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(nextNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 2,
"fromInx": currentNode.Inx,
@@ -776,6 +841,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(outNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 2,
"fromInx": currentNode.Inx,
@@ -795,6 +861,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
targetIP: target.IP,
targetPort: target.Port,
description: description,
protocol: "tcp",
metadata: map[string]interface{}{
"fromChainType": 3,
},
@@ -810,6 +877,7 @@ func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []dia
targetIP: target.IP,
targetPort: target.Port,
description: description,
protocol: "tcp",
metadata: map[string]interface{}{
"fromChainType": 1,
},
@@ -864,6 +932,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
}
ipPreference := h.repo.GetTunnelIPPreference(tunnelID)
protocol := strings.ToLower(strings.TrimSpace(tunnel.Protocol))
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
workItems := make([]diagnosisWorkItem, 0, len(chainRows)*2)
@@ -876,6 +945,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
targetIP: "www.bing.com",
targetPort: 443,
description: description,
protocol: "tcp",
metadata: map[string]interface{}{
"fromChainType": 1,
},
@@ -892,6 +962,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(firstNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 1,
"toChainType": 2,
@@ -908,6 +979,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(outNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 1,
"toChainType": 3,
@@ -928,6 +1000,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(nextNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 2,
"fromInx": currentNode.Inx,
@@ -945,6 +1018,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
hasChainHop: true,
ipPreference: ipPreference,
description: description,
protocol: defaultString(strings.ToLower(strings.TrimSpace(outNode.Protocol)), protocol),
metadata: map[string]interface{}{
"fromChainType": 2,
"fromInx": currentNode.Inx,
@@ -963,6 +1037,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
targetIP: "www.bing.com",
targetPort: 443,
description: description,
protocol: "tcp",
metadata: map[string]interface{}{
"fromChainType": 3,
},
@@ -976,6 +1051,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
targetIP: "www.bing.com",
targetPort: 443,
description: description,
protocol: "tcp",
metadata: map[string]interface{}{
"fromChainType": 1,
},
@@ -1095,9 +1171,9 @@ func (h *Handler) executeDiagnosisWorkItem(workItem diagnosisWorkItem, options d
single := make([]map[string]interface{}, 0, 1)
nodeCache := map[int64]*nodeRecord{}
if workItem.hasChainHop {
h.appendChainHopDiagnosis(&single, nodeCache, workItem.fromNodeID, workItem.toNode, workItem.description, workItem.metadata, workItem.ipPreference, options)
h.appendChainHopDiagnosis(&single, nodeCache, workItem.fromNodeID, workItem.toNode, workItem.description, workItem.metadata, workItem.ipPreference, workItem.protocol, options)
} else {
h.appendPathDiagnosis(&single, nodeCache, workItem.fromNodeID, workItem.targetIP, workItem.targetPort, workItem.description, workItem.metadata, options)
h.appendPathDiagnosis(&single, nodeCache, workItem.fromNodeID, workItem.targetIP, workItem.targetPort, workItem.description, workItem.metadata, workItem.protocol, options)
}
if len(single) == 0 {
@@ -1223,14 +1299,14 @@ func (h *Handler) appendFailedDiagnosis(results *[]map[string]interface{}, nodeC
item["nodeName"] = node.Name
}
if strings.TrimSpace(message) == "" {
message = "TCP连接失败"
message = "连接失败"
}
item["success"] = false
item["message"] = message
*results = append(*results, item)
}
func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}, options diagnosisExecOptions) {
func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, targetIP string, targetPort int, description string, metadata map[string]interface{}, protocol string, options diagnosisExecOptions) {
item := newDiagnosisResultItem(fromNodeID, targetIP, targetPort, description, metadata)
fromNode, err := h.cachedNode(nodeCache, fromNodeID)
@@ -1247,9 +1323,9 @@ func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCac
pingErr error
)
if fromNode.IsRemote == 1 {
pingData, pingErr = h.tcpPingViaRemoteNode(fromNode, targetIP, targetPort, options)
pingData, pingErr = h.pingViaRemoteNode(fromNode, targetIP, targetPort, protocol, options)
} else {
pingData, pingErr = h.tcpPingViaNode(fromNodeID, targetIP, targetPort, options)
pingData, pingErr = h.pingViaNode(fromNodeID, targetIP, targetPort, protocol, options)
}
if pingErr != nil {
item["success"] = false
@@ -1266,21 +1342,21 @@ func (h *Handler) appendPathDiagnosis(results *[]map[string]interface{}, nodeCac
message := strings.TrimSpace(asString(pingData["message"]))
if success {
if message == "" {
message = "TCP连接成功"
message = "连接成功"
}
} else {
if message == "" {
message = strings.TrimSpace(asString(pingData["errorMessage"]))
}
if message == "" {
message = "TCP连接失败"
message = "连接失败"
}
}
item["message"] = message
*results = append(*results, item)
}
func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, toNode chainNodeRecord, description string, metadata map[string]interface{}, ipPreference string, options diagnosisExecOptions) {
func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nodeCache map[int64]*nodeRecord, fromNodeID int64, toNode chainNodeRecord, description string, metadata map[string]interface{}, ipPreference string, protocol string, options diagnosisExecOptions) {
fromNode, _ := h.cachedNode(nodeCache, fromNodeID)
targetNode, err := h.cachedNode(nodeCache, toNode.NodeID)
if err != nil {
@@ -1292,7 +1368,7 @@ func (h *Handler) appendChainHopDiagnosis(results *[]map[string]interface{}, nod
h.appendFailedDiagnosis(results, nodeCache, fromNodeID, strings.Trim(strings.TrimSpace(targetNode.ServerIP), "[]"), toNode.Port, description, metadata, err.Error())
return
}
h.appendPathDiagnosis(results, nodeCache, fromNodeID, targetIP, targetPort, description, metadata, options)
h.appendPathDiagnosis(results, nodeCache, fromNodeID, targetIP, targetPort, description, metadata, protocol, options)
}
func resolveChainProbeTarget(fromNode, targetNode *nodeRecord, preferredPort int, ipPreference string, connectIp string) (string, int, error) {
@@ -1385,13 +1461,81 @@ func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int, op
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
IP: strings.TrimSpace(ip),
Port: port,
Count: 4,
Timeout: options.pingTimeoutMS,
IP: strings.TrimSpace(ip),
Port: port,
Count: 4,
Timeout: options.pingTimeoutMS,
Protocol: "tcp",
})
}
func (h *Handler) udpPingViaRemoteNode(node *nodeRecord, ip string, port int, options diagnosisExecOptions) (map[string]interface{}, error) {
if node == nil {
return nil, errors.New("节点不存在")
}
remoteURL := strings.TrimSpace(node.RemoteURL)
remoteToken := strings.TrimSpace(node.RemoteToken)
if remoteURL == "" || remoteToken == "" {
return nil, errors.New("远程节点缺少共享配置")
}
if options.commandTimeout <= 0 {
options.commandTimeout = diagnosisCommandTimeout
}
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
IP: strings.TrimSpace(ip),
Port: port,
Count: 4,
Timeout: options.pingTimeoutMS,
Protocol: "udp",
})
}
func isUDPBasedProtocol(protocol string) bool {
p := strings.ToLower(strings.TrimSpace(protocol))
return p == "kcp" || p == "udp" || p == "quic"
}
func (h *Handler) pingViaNode(nodeID int64, ip string, port int, protocol string, options diagnosisExecOptions) (map[string]interface{}, error) {
if isUDPBasedProtocol(protocol) {
return h.udpPingViaNode(nodeID, ip, port, options)
}
return h.tcpPingViaNode(nodeID, ip, port, options)
}
func (h *Handler) pingViaRemoteNode(node *nodeRecord, ip string, port int, protocol string, options diagnosisExecOptions) (map[string]interface{}, error) {
if isUDPBasedProtocol(protocol) {
return h.udpPingViaRemoteNode(node, ip, port, options)
}
return h.tcpPingViaRemoteNode(node, ip, port, options)
}
func (h *Handler) udpPingViaNode(nodeID int64, ip string, port int, options diagnosisExecOptions) (map[string]interface{}, error) {
if options.commandTimeout <= 0 {
options.commandTimeout = diagnosisCommandTimeout
}
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
res, err := h.sendNodeCommandWithTimeout(nodeID, "UdpPing", map[string]interface{}{
"ip": ip,
"port": port,
"count": 4,
"timeout": options.pingTimeoutMS,
}, options.commandTimeout, false, false)
if err != nil {
return nil, err
}
if res.Data == nil {
return nil, errors.New("节点未返回诊断数据")
}
return res.Data, nil
}
func splitRemoteTargets(remoteAddr string) []string {
parts := strings.Split(remoteAddr, ",")
out := make([]string, 0, len(parts))
@@ -1549,7 +1693,7 @@ func compactErrorMessage(msg string) string {
return strings.Join(strings.Fields(strings.ToLower(msg)), "")
}
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64) []map[string]interface{} {
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, runtimeLimiters forwardRuntimeLimiters) []map[string]interface{} {
protocols := []string{"tcp", "udp"}
services := make([]map[string]interface{}, 0, 2)
targets := splitRemoteTargets(forward.RemoteAddr)
@@ -1592,6 +1736,19 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
},
},
}
if runtimeLimiters.ConnLimiter != "" {
service["climiter"] = runtimeLimiters.ConnLimiter
}
if runtimeLimiters.TrafficLimiter != "" {
service["limiter"] = runtimeLimiters.TrafficLimiter
}
if forward.ProxyProtocol > 0 {
handlerConfig := service["handler"].(map[string]interface{})
if handlerConfig["metadata"] == nil {
handlerConfig["metadata"] = map[string]interface{}{}
}
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
}
if protocol == "udp" {
listenerMetadata := map[string]interface{}{
"keepAlive": true,
@@ -1603,10 +1760,10 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
}
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
}
if limiterID != nil && *limiterID > 0 {
service["limiter"] = strconv.FormatInt(*limiterID, 10)
if service["metadata"] == nil {
service["metadata"] = map[string]interface{}{}
}
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
}
services = append(services, service)
}
@@ -1707,14 +1864,76 @@ func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int)
return nil
}
func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) {
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, cfg forwardLimiterConfig) error {
if cfg.Name == "" || len(cfg.Limits) == 0 {
return nil
}
payload := map[string]interface{}{"name": cfg.Name, "limits": cfg.Limits}
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
if !isAlreadyExistsMessage(err.Error()) {
return fmt.Errorf("连接限制器下发失败: %w", err)
}
updatePayload := map[string]interface{}{"limiter": cfg.Name, "data": payload}
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
}
}
return nil
}
func buildConnLimiterConfigs(forward *forwardRecord, userMaxConn int) []forwardLimiterConfig {
if forward == nil {
return nil
}
if forward.MaxConn > 0 {
limits := []string{fmt.Sprintf("$ %d", forward.MaxConn)}
if forward.IPMaxConn > 0 {
limits = append(limits, fmt.Sprintf("$$ %d", forward.IPMaxConn))
}
return []forwardLimiterConfig{{Name: fmt.Sprintf("rule_conn_limit_%d", forward.ID), Limits: limits}}
}
configs := make([]forwardLimiterConfig, 0, 2)
if userMaxConn > 0 {
configs = append(configs, forwardLimiterConfig{Name: fmt.Sprintf("user_conn_limit_%d", forward.UserID), Limits: []string{fmt.Sprintf("$ %d", userMaxConn)}})
}
if forward.IPMaxConn > 0 {
configs = append(configs, forwardLimiterConfig{Name: fmt.Sprintf("rule_conn_limit_%d", forward.ID), Limits: []string{fmt.Sprintf("$$ %d", forward.IPMaxConn)}})
}
return configs
}
func joinLimiterNames(configs []forwardLimiterConfig) string {
names := make([]string, 0, len(configs))
for _, cfg := range configs {
if cfg.Name != "" {
names = append(names, cfg.Name)
}
}
return strings.Join(names, ",")
}
func speedToLimitLine(key string, speed int) string {
rate := float64(speed) / 8.0
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
return fmt.Sprintf("%s %.1fMB %.1fMB", key, rate, rate)
}
func buildTrafficLimiterPayload(name string, totalSpeed *int, ipSpeed *int) map[string]interface{} {
limits := make([]string, 0, 3)
if totalSpeed != nil && *totalSpeed > 0 {
limits = append(limits, speedToLimitLine("$", *totalSpeed))
}
if ipSpeed != nil && *ipSpeed > 0 {
limits = append(limits, speedToLimitLine("0.0.0.0/0", *ipSpeed), speedToLimitLine("::/0", *ipSpeed))
}
return map[string]interface{}{"name": name, "limits": limits}
}
func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) {
name := strconv.FormatInt(limiterID, 10)
return name, map[string]interface{}{
"name": name,
"limits": []string{limitStr},
"limits": []string{speedToLimitLine("$", speed)},
}
}
@@ -1742,3 +1961,20 @@ func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int)
return nil
}
func (h *Handler) ensureTrafficLimiterOnNode(nodeID int64, name string, totalSpeed *int, ipSpeed *int) error {
payload := buildTrafficLimiterPayload(name, totalSpeed, ipSpeed)
limits, _ := payload["limits"].([]string)
if name == "" || len(limits) == 0 {
return nil
}
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
if !isAlreadyExistsMessage(err.Error()) {
return fmt.Errorf("限速规则下发失败: %w", err)
}
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil {
return fmt.Errorf("限速规则更新失败: %w", updateErr)
}
}
return nil
}
@@ -378,7 +378,7 @@ func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) {
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil)
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", forwardRuntimeLimiters{})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
@@ -393,7 +393,7 @@ func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil)
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
@@ -409,7 +409,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil)
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", forwardRuntimeLimiters{})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
@@ -464,7 +464,7 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
t.Run(tt.name, func(t *testing.T) {
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil)
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, forwardRuntimeLimiters{})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
@@ -478,6 +478,58 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
}
}
func TestBuildConnLimiterConfigCombinesTotalAndPerIP(t *testing.T) {
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, MaxConn: 100, IPMaxConn: 5}, 37)
want := []forwardLimiterConfig{{Name: "rule_conn_limit_42", Limits: []string{"$ 100", "$$ 5"}}}
if !reflect.DeepEqual(cfgs, want) {
t.Fatalf("expected %+v, got %+v", want, cfgs)
}
}
func TestBuildConnLimiterConfigUsesUserTotalWithRulePerIP(t *testing.T) {
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, IPMaxConn: 5}, 37)
want := []forwardLimiterConfig{
{Name: "user_conn_limit_9", Limits: []string{"$ 37"}},
{Name: "rule_conn_limit_42", Limits: []string{"$$ 5"}},
}
if !reflect.DeepEqual(cfgs, want) {
t.Fatalf("expected %+v, got %+v", want, cfgs)
}
if got := joinLimiterNames(cfgs); got != "user_conn_limit_9,rule_conn_limit_42" {
t.Fatalf("expected composite limiter names, got %q", got)
}
}
func TestBuildTrafficLimiterPayloadUsesOnlyPerIPRulesWhenTotalIsSeparate(t *testing.T) {
payload := buildTrafficLimiterPayload("rule_traffic_limit_42", nil, intPtr(40))
wantLimits := []string{"0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
if payload["name"] != "rule_traffic_limit_42" {
t.Fatalf("expected name rule_traffic_limit_42, got %v", payload["name"])
}
if !reflect.DeepEqual(payload["limits"], wantLimits) {
t.Fatalf("expected limits %v, got %v", wantLimits, payload["limits"])
}
}
func TestBuildForwardServiceConfigsUsesRuntimeLimiterNames(t *testing.T) {
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{TrafficLimiter: "rule_traffic_limit_42", ConnLimiter: "rule_conn_limit_42"})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
for _, service := range services {
if service["limiter"] != "rule_traffic_limit_42" {
t.Fatalf("expected traffic limiter rule_traffic_limit_42, got %v", service["limiter"])
}
if service["climiter"] != "rule_conn_limit_42" {
t.Fatalf("expected conn limiter rule_conn_limit_42, got %v", service["climiter"])
}
}
}
func intPtr(v int) *int { return &v }
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
tests := []struct {
name string
+12 -6
View File
@@ -86,10 +86,11 @@ type federationRuntimeReleaseRoleRequest struct {
}
type federationRuntimeDiagnoseRequest struct {
IP string `json:"ip"`
Port int `json:"port"`
Count int `json:"count"`
Timeout int `json:"timeout"`
IP string `json:"ip"`
Port int `json:"port"`
Count int `json:"count"`
Timeout int `json:"timeout"`
Protocol string `json:"protocol"`
}
type federationRuntimeCommandRequest struct {
@@ -649,7 +650,7 @@ func (h *Handler) nodeImport(w http.ResponseWriter, r *http.Request) {
return
}
if err := IsSafeRemoteAddr(rURL.Host); err != nil {
response.WriteJSON(w, response.Err(403, "禁止将远程节点地址设置为内部网络或保留地址"))
response.WriteJSON(w, response.Err(403, "禁止将远程节点地址设置为内部网络"))
return
}
@@ -1281,7 +1282,12 @@ func (h *Handler) federationRuntimeDiagnose(w http.ResponseWriter, r *http.Reque
commandTimeout = diagnosisCommandTimeout
}
res, err := h.sendNodeCommandWithTimeout(share.NodeID, "TcpPing", map[string]interface{}{
commandType := "TcpPing"
if isUDPBasedProtocol(req.Protocol) {
commandType = "UdpPing"
}
res, err := h.sendNodeCommandWithTimeout(share.NodeID, commandType, map[string]interface{}{
"ip": req.IP,
"port": req.Port,
"count": req.Count,
@@ -43,16 +43,19 @@ func (h *Handler) processFlowItem(nodeID int64, item flowItem) {
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
if ok {
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil {
h.enforceUserQuotaIfNeeded(userID, quota)
if h.forwardExists(forwardID) {
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil {
h.enforceUserQuotaIfNeeded(userID, quota)
}
if userTunnelID > 0 {
h.enforceFlowPolicies(userID, userTunnelID)
}
} else if nodeID > 0 {
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
}
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
if userTunnelID > 0 {
h.enforceFlowPolicies(userID, userTunnelID)
}
return
}
@@ -553,6 +556,14 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem
}
parts := strings.Split(name, "_")
if len(parts) == 2 && parts[0] == "tunnel" {
tunnelID, err := strconv.ParseInt(parts[1], 10, 64)
if err == nil && tunnelID > 0 && !h.tunnelExists(tunnelID) {
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
}
continue
}
if len(parts) >= 3 {
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
if err == nil && forwardID > 0 && hasUnboundForwardPeerRuntime {
@@ -566,7 +577,7 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem
suffix := parts[len(parts)-1]
switch suffix {
case "tls":
case "tls", "kcp", "wss", "mtls", "mwss", "mtcp":
tunnelID, err := strconv.ParseInt(parts[0], 10, 64)
if err != nil || tunnelID <= 0 || h.tunnelExists(tunnelID) {
continue
@@ -574,6 +585,10 @@ func (h *Handler) cleanOrphanedServices(nodeID int64, services []namedConfigItem
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
case "tcp":
if len(parts) < 4 {
tunnelID, err := strconv.ParseInt(parts[0], 10, 64)
if err == nil && tunnelID > 0 && !h.tunnelExists(tunnelID) {
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{name}}, false, true)
}
continue
}
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
@@ -628,6 +643,21 @@ func (h *Handler) forwardExists(forwardID int64) bool {
return ok
}
func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName string) {
parts := strings.Split(serviceName, "_")
if len(parts) < 3 {
return
}
forwardID, err := strconv.ParseInt(parts[0], 10, 64)
if err != nil || forwardID <= 0 {
return
}
base := parts[0] + "_" + parts[1] + "_" + parts[2]
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{
"services": []string{base + "_tcp", base + "_udp"},
}, false, true)
}
func (h *Handler) speedLimiterExists(name string) bool {
if name == "" {
return false
@@ -0,0 +1,183 @@
package handler
import (
"log"
"sort"
"strings"
"time"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
)
type flowPolicyTarget struct {
UserID int64
UserTunnelID int64
}
type flowUploadBatch struct {
flowDeltas []repo.FlowUploadCounterDelta
quotaUsage map[int64]int64
policyTargets []flowPolicyTarget
forwardTraffic map[int64]tunnelTrafficDelta
orphanServices map[string]struct{}
peerShareForwardItems map[string]flowItem
peerShareRuntimeItems map[int64]flowItem
}
func (h *Handler) buildFlowUploadBatch(items []flowItem, metas map[int64]repo.FlowUploadForwardMeta) flowUploadBatch {
batch := flowUploadBatch{
quotaUsage: make(map[int64]int64),
forwardTraffic: make(map[int64]tunnelTrafficDelta),
orphanServices: make(map[string]struct{}),
peerShareForwardItems: make(map[string]flowItem),
peerShareRuntimeItems: make(map[int64]flowItem),
}
policySeen := map[flowPolicyTarget]struct{}{}
flowSeen := map[int64]int{}
for _, item := range items {
serviceName := strings.TrimSpace(item.N)
if serviceName == "" || serviceName == "web_api" {
continue
}
if runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName); ok {
merged := batch.peerShareRuntimeItems[runtimeID]
merged.N = serviceName
merged.U += item.U
merged.D += item.D
batch.peerShareRuntimeItems[runtimeID] = merged
continue
}
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
if !ok {
continue
}
normalized := normalizeForwardRuntimeServiceName(serviceName)
merged := batch.peerShareForwardItems[normalized]
merged.N = normalized
merged.U += item.U
merged.D += item.D
batch.peerShareForwardItems[normalized] = merged
meta, exists := metas[forwardID]
if !exists {
batch.orphanServices[serviceName] = struct{}{}
continue
}
raw := batch.forwardTraffic[forwardID]
raw.bytesIn += item.D
raw.bytesOut += item.U
batch.forwardTraffic[forwardID] = raw
scaledIn := int64(float64(item.D)*meta.TrafficRatio) * meta.TunnelFlow
scaledOut := int64(float64(item.U)*meta.TrafficRatio) * meta.TunnelFlow
if idx, ok := flowSeen[forwardID]; ok {
batch.flowDeltas[idx].InFlow += scaledIn
batch.flowDeltas[idx].OutFlow += scaledOut
} else {
flowSeen[forwardID] = len(batch.flowDeltas)
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{
ForwardID: forwardID,
UserID: userID,
UserTunnelID: userTunnelID,
InFlow: scaledIn,
OutFlow: scaledOut,
})
}
batch.quotaUsage[userID] += scaledIn + scaledOut
target := flowPolicyTarget{UserID: userID, UserTunnelID: userTunnelID}
if _, seen := policySeen[target]; !seen {
policySeen[target] = struct{}{}
batch.policyTargets = append(batch.policyTargets, target)
}
}
sort.Slice(batch.policyTargets, func(i, j int) bool {
if batch.policyTargets[i].UserID == batch.policyTargets[j].UserID {
return batch.policyTargets[i].UserTunnelID < batch.policyTargets[j].UserTunnelID
}
return batch.policyTargets[i].UserID < batch.policyTargets[j].UserID
})
return batch
}
func (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now time.Time) {
if h == nil || h.repo == nil {
return
}
h.applyFlowDeltasWithFallback(nodeID, batch.flowDeltas)
for userID, quota := range h.applyQuotaUsageWithFallback(nodeID, batch.quotaUsage, now) {
h.enforceUserQuotaIfNeeded(userID, quota)
}
for _, target := range batch.policyTargets {
if target.UserID <= 0 || target.UserTunnelID <= 0 {
continue
}
h.enforceFlowPolicies(target.UserID, target.UserTunnelID)
}
for serviceName := range batch.orphanServices {
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
}
for serviceName, item := range batch.peerShareForwardItems {
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
if ok {
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
}
}
for runtimeID, item := range batch.peerShareRuntimeItems {
h.processPeerShareFlow(runtimeID, item)
}
}
func (h *Handler) applyFlowDeltasWithFallback(nodeID int64, deltas []repo.FlowUploadCounterDelta) {
if h == nil || h.repo == nil || len(deltas) == 0 {
return
}
if err := h.repo.ApplyFlowUploadDeltasBatch(deltas); err == nil {
return
} else {
log.Printf("flow upload write failed op=flow.batch_apply node_id=%d err=%v", nodeID, err)
}
for _, delta := range deltas {
if err := h.repo.AddFlow(delta.ForwardID, delta.UserID, delta.UserTunnelID, delta.InFlow, delta.OutFlow); err != nil {
log.Printf("flow upload write failed op=flow.single_apply node_id=%d forward_id=%d user_id=%d user_tunnel_id=%d err=%v", nodeID, delta.ForwardID, delta.UserID, delta.UserTunnelID, err)
}
}
}
func (h *Handler) applyQuotaUsageWithFallback(nodeID int64, usages map[int64]int64, now time.Time) map[int64]*model.UserQuotaView {
if h == nil || h.repo == nil || len(usages) == 0 {
return map[int64]*model.UserQuotaView{}
}
quotaViews, err := h.repo.AddUserQuotaUsageBatch(usages, now)
if err == nil {
return quotaViews
}
log.Printf("flow upload write failed op=quota.batch_apply node_id=%d err=%v", nodeID, err)
userIDs := make([]int64, 0, len(usages))
for userID := range usages {
if userID > 0 {
userIDs = append(userIDs, userID)
}
}
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
quotaViews = make(map[int64]*model.UserQuotaView, len(userIDs))
for _, userID := range userIDs {
quota, singleErr := h.repo.AddUserQuotaUsage(userID, usages[userID], now)
if singleErr != nil {
log.Printf("flow upload write failed op=quota.single_apply node_id=%d user_id=%d err=%v", nodeID, userID, singleErr)
continue
}
if quota != nil {
quotaViews[userID] = quota
}
}
return quotaViews
}
@@ -0,0 +1,254 @@
package handler
import (
"path/filepath"
"testing"
"time"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
)
func TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets(t *testing.T) {
h := &Handler{}
metas := map[int64]repo.FlowUploadForwardMeta{
20: {
ForwardID: 20,
TunnelID: 1,
TrafficRatio: 2,
TunnelFlow: 3,
},
}
batch := h.buildFlowUploadBatch([]flowItem{
{N: "20_2_10", U: 70, D: 50},
{N: "20_2_10_tcp", U: 40, D: 30},
{N: "99_2_10", U: 12, D: 8},
{N: "fed_svc_17", U: 9, D: 1},
}, metas)
if len(batch.flowDeltas) != 1 {
t.Fatalf("expected 1 flow delta, got %d", len(batch.flowDeltas))
}
delta := batch.flowDeltas[0]
if delta.ForwardID != 20 || delta.UserID != 2 || delta.UserTunnelID != 10 {
t.Fatalf("unexpected flow delta identity: %#v", delta)
}
if delta.InFlow != 480 || delta.OutFlow != 660 {
t.Fatalf("expected scaled flow in=480 out=660, got in=%d out=%d", delta.InFlow, delta.OutFlow)
}
if batch.quotaUsage[2] != 1140 {
t.Fatalf("expected quota usage 1140, got %d", batch.quotaUsage[2])
}
if len(batch.policyTargets) != 1 {
t.Fatalf("expected 1 policy target, got %d", len(batch.policyTargets))
}
if batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
t.Fatalf("unexpected policy target: %#v", batch.policyTargets[0])
}
traffic := batch.forwardTraffic[20]
if traffic.bytesIn != 80 || traffic.bytesOut != 110 {
t.Fatalf("expected raw traffic in=80 out=110, got in=%d out=%d", traffic.bytesIn, traffic.bytesOut)
}
if _, ok := batch.orphanServices["99_2_10"]; !ok {
t.Fatalf("expected orphan service cleanup target for 99_2_10")
}
if item, ok := batch.peerShareForwardItems["99_2_10"]; !ok || item.U != 12 || item.D != 8 {
t.Fatalf("expected orphan forward to remain eligible for peer-share accounting, got %#v ok=%v", item, ok)
}
if item, ok := batch.peerShareForwardItems["20_2_10"]; !ok || item.U != 110 || item.D != 80 {
t.Fatalf("expected merged peer-share forward item, got %#v ok=%v", item, ok)
}
if item, ok := batch.peerShareRuntimeItems[17]; !ok || item.U != 9 || item.D != 1 {
t.Fatalf("expected merged peer-share runtime item, got %#v ok=%v", item, ok)
}
}
func TestApplyFlowUploadBatchContinuesPolicyAndPeerShareSideEffectsWhenQuotaBatchFails(t *testing.T) {
r, err := repo.Open(filepath.Join(t.TempDir(), "flow-upload-batch-quota-fail.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
if err := r.DB().Create(&model.User{ID: 2, User: "flow-user", Pwd: "pwd", RoleID: 1, ExpTime: 2727251700000, Flow: 99999, Num: 99999, CreatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed user: %v", err)
}
if err := r.DB().Create(&model.Tunnel{ID: 1, Name: "tunnel-1", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed tunnel: %v", err)
}
if err := r.DB().Create(&model.UserTunnel{ID: 10, UserID: 2, TunnelID: 1, Num: 99999, Flow: 0, ExpTime: 2727251700000, Status: 1}).Error; err != nil {
t.Fatalf("seed user tunnel: %v", err)
}
if err := r.DB().Create(&model.Forward{ID: 20, UserID: 2, UserName: "flow-user", Name: "forward-20", TunnelID: 1, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed forward: %v", err)
}
if err := r.CreatePeerShare(&repo.PeerShare{Name: "share", NodeID: 1, Token: "token", MaxBandwidth: 0, CurrentFlow: 0, PortRangeStart: 31000, PortRangeEnd: 31010, IsActive: 1, CreatedTime: nowMs, UpdatedTime: nowMs}); err != nil {
t.Fatalf("create peer share: %v", err)
}
share, err := r.GetPeerShareByToken("token")
if err != nil || share == nil {
t.Fatalf("load peer share: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, share.ID, 1, "svc-r1", "svc-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31001, "", 1, 1, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert peer share runtime: %v", err)
}
if err := r.DB().Exec(`
CREATE TRIGGER fail_user_quota_insert
BEFORE INSERT ON user_quota
BEGIN
SELECT RAISE(FAIL, 'quota insert blocked for test');
END;
`).Error; err != nil {
t.Fatalf("create quota failure trigger: %v", err)
}
h := &Handler{repo: r}
h.applyFlowUploadBatch(1, flowUploadBatch{
flowDeltas: []repo.FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 80, OutFlow: 120}},
quotaUsage: map[int64]int64{2: 200},
policyTargets: []flowPolicyTarget{{UserID: 2, UserTunnelID: 10}},
peerShareForwardItems: map[string]flowItem{"20_2_10": {N: "20_2_10", U: 120, D: 80}},
}, now)
if got := mustQueryInt(t, r, `SELECT status FROM forward WHERE id = 20`); got != 0 {
t.Fatalf("expected flow-policy enforcement to pause forward after quota failure, got status=%d", got)
}
updatedShare, err := r.GetPeerShare(share.ID)
if err != nil || updatedShare == nil {
t.Fatalf("reload peer share: %v", err)
}
if updatedShare.CurrentFlow != 200 {
t.Fatalf("expected peer-share flow accounting to continue after quota failure, got %d", updatedShare.CurrentFlow)
}
}
func TestApplyFlowUploadBatchContinuesPeerShareSideEffectsWhenFlowBatchFails(t *testing.T) {
r, err := repo.Open(filepath.Join(t.TempDir(), "flow-upload-batch-flow-fail.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
if err := r.DB().Create(&model.User{ID: 2, User: "flow-user", Pwd: "pwd", RoleID: 1, ExpTime: 2727251700000, Flow: 99999, Num: 99999, CreatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed user: %v", err)
}
if err := r.DB().Create(&model.Tunnel{ID: 1, Name: "tunnel-1", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed tunnel: %v", err)
}
if err := r.DB().Create(&model.UserTunnel{ID: 10, UserID: 2, TunnelID: 1, Num: 99999, Flow: 0, ExpTime: 2727251700000, Status: 1}).Error; err != nil {
t.Fatalf("seed user tunnel: %v", err)
}
if err := r.DB().Create(&model.Forward{ID: 20, UserID: 2, UserName: "flow-user", Name: "forward-20", TunnelID: 1, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed forward: %v", err)
}
if err := r.DB().Create(&model.Forward{ID: 21, UserID: 2, UserName: "flow-user", Name: "forward-21", TunnelID: 1, RemoteAddr: "1.1.1.1:81", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}).Error; err != nil {
t.Fatalf("seed second forward: %v", err)
}
if err := r.CreatePeerShare(&repo.PeerShare{Name: "share", NodeID: 1, Token: "token", MaxBandwidth: 0, CurrentFlow: 0, PortRangeStart: 31000, PortRangeEnd: 31010, IsActive: 1, CreatedTime: nowMs, UpdatedTime: nowMs}); err != nil {
t.Fatalf("create peer share: %v", err)
}
share, err := r.GetPeerShareByToken("token")
if err != nil || share == nil {
t.Fatalf("load peer share: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, share.ID, 1, "svc-r1", "svc-rk1", "", "forward", "", "20_2_10", "tcp", "fifo", 31001, "", 1, 1, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert peer share runtime: %v", err)
}
if err := r.DB().Exec(`
CREATE TRIGGER fail_forward_flow_update
BEFORE UPDATE ON forward
WHEN NEW.id = 21 AND (NEW.in_flow != OLD.in_flow OR NEW.out_flow != OLD.out_flow)
BEGIN
SELECT RAISE(FAIL, 'forward flow update blocked for test');
END;
`).Error; err != nil {
t.Fatalf("create flow failure trigger: %v", err)
}
h := &Handler{repo: r}
h.applyFlowUploadBatch(1, flowUploadBatch{
flowDeltas: []repo.FlowUploadCounterDelta{
{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 80, OutFlow: 120},
{ForwardID: 21, UserID: 2, UserTunnelID: 10, InFlow: 30, OutFlow: 40},
},
quotaUsage: map[int64]int64{2: 200},
policyTargets: []flowPolicyTarget{{UserID: 2, UserTunnelID: 10}},
peerShareForwardItems: map[string]flowItem{"20_2_10": {N: "20_2_10", U: 120, D: 80}},
}, now)
if got := mustQueryInt(t, r, `SELECT status FROM forward WHERE id = 20`); got != 0 {
t.Fatalf("expected flow-policy enforcement to pause forward after flow batch failure, got status=%d", got)
}
updatedShare, err := r.GetPeerShare(share.ID)
if err != nil || updatedShare == nil {
t.Fatalf("reload peer share: %v", err)
}
if updatedShare.CurrentFlow != 200 {
t.Fatalf("expected peer-share flow accounting to continue after flow batch failure, got %d", updatedShare.CurrentFlow)
}
if got := mustQueryInt(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 80 {
t.Fatalf("expected flow fallback to persist forward 20 in_flow=80, got %d", got)
}
if got := mustQueryInt(t, r, `SELECT in_flow FROM forward WHERE id = 21`); got != 0 {
t.Fatalf("expected failed forward 21 delta to remain unapplied, got %d", got)
}
if got := mustQueryInt(t, r, `SELECT in_flow FROM user WHERE id = 2`); got != 80 {
t.Fatalf("expected flow fallback to preserve successful user totals, got %d", got)
}
if got := mustQueryInt(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 80 {
t.Fatalf("expected flow fallback to preserve successful user_tunnel totals, got %d", got)
}
}
func TestApplyFlowUploadBatchFallsBackToPerUserQuotaUpdates(t *testing.T) {
r, err := repo.Open(filepath.Join(t.TempDir(), "flow-upload-batch-quota-fallback.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
monthKey := int64(now.Year()*100 + int(now.Month()))
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user 2: %v", err)
}
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(3, 'u3', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user 3: %v", err)
}
if err := r.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 0, 0, 0, 0, ?, ?, 0, 0, '', ?, ?), (3, 0, 0, 0, 0, ?, ?, 0, 0, '', ?, ?)`, dayKey, monthKey, nowMs, nowMs, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user quotas: %v", err)
}
if err := r.DB().Exec(`
CREATE TRIGGER fail_user_3_quota_update
BEFORE UPDATE ON user_quota
WHEN NEW.user_id = 3 AND (NEW.daily_used_bytes != OLD.daily_used_bytes OR NEW.monthly_used_bytes != OLD.monthly_used_bytes)
BEGIN
SELECT RAISE(FAIL, 'quota update blocked for user 3');
END;
`).Error; err != nil {
t.Fatalf("create quota fallback trigger: %v", err)
}
h := &Handler{repo: r}
h.applyFlowUploadBatch(1, flowUploadBatch{quotaUsage: map[int64]int64{2: 200, 3: 300}}, now)
if got := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 2`); got != 200 {
t.Fatalf("expected quota fallback to persist user 2 usage, got %d", got)
}
if got := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 3`); got != 0 {
t.Fatalf("expected failed user 3 quota delta to remain unapplied, got %d", got)
}
}
@@ -0,0 +1,123 @@
package handler
import (
"database/sql"
"testing"
"time"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
)
func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing.T) {
forward := &forwardRecord{
ID: 1,
UserID: 2,
TunnelID: 3,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
ProxyProtocol: 2,
}
tunnel := &tunnelRecord{Type: 1}
node := &nodeRecord{
InterfaceName: "eth0",
TCPListenAddr: "0.0.0.0",
UDPListenAddr: "0.0.0.0",
}
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{})
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
for _, service := range services {
serviceMetadata, ok := service["metadata"].(map[string]interface{})
if !ok {
t.Fatalf("expected metadata map, got %T", service["metadata"])
}
if serviceMetadata["interface"] != "eth0" {
t.Fatalf("expected interface metadata eth0, got %v", serviceMetadata["interface"])
}
if _, ok := serviceMetadata["proxyProtocol"]; ok {
t.Fatalf("proxyProtocol should not be listener metadata: %v", serviceMetadata)
}
handlerConfig, ok := service["handler"].(map[string]interface{})
if !ok {
t.Fatalf("expected handler config map, got %T", service["handler"])
}
handlerMetadata, ok := handlerConfig["metadata"].(map[string]interface{})
if !ok {
t.Fatalf("expected handler metadata map, got %T", handlerConfig["metadata"])
}
if handlerMetadata["proxyProtocol"] != 2 {
t.Fatalf("expected handler proxyProtocol 2, got %v", handlerMetadata["proxyProtocol"])
}
}
}
func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
r, err := repo.Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Create(&model.Forward{
UserID: 2,
UserName: "rollback-user",
Name: "rollback-forward",
TunnelID: 3,
RemoteAddr: "9.9.9.9:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
IPMaxConn: 5,
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
ProxyProtocol: 2,
}).Error; err != nil {
t.Fatalf("create forward: %v", err)
}
forwardID := mustLastInsertID(t, r, "rollback-forward")
if err := r.DB().Model(&model.Forward{}).Where("id = ?", forwardID).Updates(map[string]interface{}{
"name": "changed-forward",
"ip_max_conn": 0,
"ip_speed_id": nil,
"proxy_protocol": 0,
"updated_time": now + 1,
}).Error; err != nil {
t.Fatalf("mutate forward: %v", err)
}
h := &Handler{repo: r}
h.rollbackForwardMutation(&forwardRecord{
ID: forwardID,
UserID: 2,
UserName: "rollback-user",
Name: "rollback-forward",
TunnelID: 3,
RemoteAddr: "9.9.9.9:443",
Strategy: "fifo",
Status: 1,
IPMaxConn: 5,
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
ProxyProtocol: 2,
}, nil)
var record model.Forward
if err := r.DB().Where("id = ?", forwardID).First(&record).Error; err != nil {
t.Fatalf("query forward: %v", err)
}
if record.ProxyProtocol != 2 {
t.Fatalf("expected proxyProtocol restored to 2, got %d", record.ProxyProtocol)
}
if record.IPMaxConn != 5 {
t.Fatalf("expected ipMaxConn restored to 5, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
t.Fatalf("expected ipSpeedId restored to 21, got %+v", record.IPSpeedID)
}
}
+39 -13
View File
@@ -7,6 +7,7 @@ import (
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"net/url"
"sort"
@@ -43,13 +44,17 @@ type Handler struct {
jobsStarted bool
jobsWG sync.WaitGroup
upgradeMu sync.Mutex
pendingUpgradeRedeploy map[int64]struct{}
upgradeMu sync.Mutex
pendingUpgradeRedeploy map[int64]struct{}
nodeOnlineRedeployAt map[int64]time.Time
nodeOnlineRedeployQueued map[int64]struct{}
nodeOnlineRedeploying map[int64]struct{}
qualityProber *tunnelQualityProber
}
const monitorTunnelQualityEnabledConfigKey = "monitor_tunnel_quality_enabled"
const allowLocalRemoteAddrConfigKey = "allow_local_remote_addr"
type loginRequest struct {
Username string `json:"username"`
@@ -95,13 +100,16 @@ const (
func New(repo *repo.Repository, jwtSecret string) *Handler {
h := &Handler{
repo: repo,
jwtSecret: jwtSecret,
wsServer: ws.NewServer(repo, jwtSecret),
metrics: metrics.NewIngestionService(repo),
healthCheck: nil,
captchaTokens: make(map[string]int64),
pendingUpgradeRedeploy: make(map[int64]struct{}),
repo: repo,
jwtSecret: jwtSecret,
wsServer: ws.NewServer(repo, jwtSecret),
metrics: metrics.NewIngestionService(repo),
healthCheck: nil,
captchaTokens: make(map[string]int64),
pendingUpgradeRedeploy: make(map[int64]struct{}),
nodeOnlineRedeployAt: make(map[int64]time.Time),
nodeOnlineRedeployQueued: make(map[int64]struct{}),
nodeOnlineRedeploying: make(map[int64]struct{}),
}
h.healthCheck = health.NewChecker(repo, h.wsServer)
h.qualityProber = newTunnelQualityProber(h)
@@ -796,11 +804,16 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
if err == nil && strings.TrimSpace(raw) != "" {
var items []flowItem
if json.Unmarshal([]byte(raw), &items) == nil {
nowMs := time.Now().UnixMilli()
h.recordTunnelMetricsFromFlowItems(node.ID, items, nowMs)
for _, item := range items {
h.processFlowItem(node.ID, item)
now := time.Now()
forwardIDs := collectFlowUploadForwardIDs(items)
metas, metaErr := h.repo.GetFlowUploadForwardMetas(forwardIDs)
if metaErr != nil {
log.Printf("flow upload metadata lookup failed node_id=%d err=%v", node.ID, metaErr)
metas = map[int64]repo.FlowUploadForwardMeta{}
}
batch := h.buildFlowUploadBatch(items, metas)
h.recordTunnelMetricsFromForwardBatch(node.ID, batch.forwardTraffic, metas, now.UnixMilli())
h.applyFlowUploadBatch(node.ID, batch, now)
}
}
@@ -1041,6 +1054,19 @@ func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
return strings.TrimSpace(strings.ToLower(cfg.Value)) != "false"
}
func (h *Handler) allowLocalRemoteAddr() bool {
if h == nil || h.repo == nil {
return false
}
cfg, err := h.repo.GetConfigByName(allowLocalRemoteAddrConfigKey)
if err != nil || cfg == nil {
return false
}
return strings.TrimSpace(strings.ToLower(cfg.Value)) == "true"
}
func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
+1 -1
View File
@@ -121,7 +121,7 @@ func (h *Handler) runHealthChecks(ctx context.Context) {
func (h *Handler) runTunnelQualityProber(ctx context.Context) {
defer h.jobsWG.Done()
if h == nil || h.qualityProber == nil || !h.isTunnelQualityMonitoringEnabled() {
if h == nil || h.qualityProber == nil {
return
}
+194 -24
View File
@@ -68,8 +68,9 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
}
roleID := 1
now := time.Now().UnixMilli()
maxConn := asInt(req["maxConn"], 0)
userID, err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, now)
userID, err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -137,6 +138,15 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("请不要作死"))
return
}
oldUser, err := h.repo.GetUserByID(id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if oldUser == nil {
response.WriteJSON(w, response.ErrDefault("用户不存在"))
return
}
dup, err := h.repo.UserExistsExcluding(username, id)
if err != nil {
@@ -156,15 +166,16 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
_, hasDailyQuota := req["dailyQuotaGB"]
_, hasMonthlyQuota := req["monthlyQuotaGB"]
now := time.Now().UnixMilli()
maxConn := asInt(req["maxConn"], 0)
pwd := asString(req["pwd"])
if strings.TrimSpace(pwd) == "" {
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, now); err != nil {
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
} else {
if err := h.repo.UpdateUserWithPassword(id, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, now); err != nil {
if err := h.repo.UpdateUserWithPassword(id, username, security.MD5(pwd), flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -208,6 +219,17 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
}
}
}
if oldUser.MaxConn != maxConn {
warnings, syncErr := h.syncUserMaxConnForwards(id)
if syncErr != nil {
response.WriteJSON(w, response.ErrDefault(fmt.Sprintf("最大连接数下发失败: %v", syncErr)))
return
}
if len(warnings) > 0 {
response.WriteJSON(w, response.OK(map[string]interface{}{"warnings": warnings}))
return
}
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -657,11 +679,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
if trimmed := strings.TrimSpace(inIP); trimmed != "" {
tunnelInIP = sql.NullString{String: trimmed, Valid: true}
}
tunnelProtocol := "tls"
if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
tunnelProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
}
tunnel := model.Tunnel{
Name: name,
TrafficRatio: trafficRatio,
Type: typeVal,
Protocol: "tls",
Protocol: tunnelProtocol,
Flow: flow,
CreatedTime: now,
UpdatedTime: now,
@@ -702,7 +728,7 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
if typeVal == 2 {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil {
h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID)
h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID, tunnelProtocol)
h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.deleteTunnelByID(tunnelID)
response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
@@ -722,17 +748,30 @@ func (h *Handler) cleanupTunnelRuntime(tunnelID int64) {
return
}
serviceName := fmt.Sprintf("%d_tls", tunnelID)
protocol := strings.TrimSpace(tunnel.Protocol)
if protocol == "" {
protocol = "tls"
}
chainName := fmt.Sprintf("chains_%d", tunnelID)
serviceNames := []string{
fmt.Sprintf("tunnel_%d", tunnelID),
fmt.Sprintf("%d_tls", tunnelID),
fmt.Sprintf("%d_kcp", tunnelID),
fmt.Sprintf("%d_wss", tunnelID),
fmt.Sprintf("%d_mtls", tunnelID),
fmt.Sprintf("%d_mwss", tunnelID),
fmt.Sprintf("%d_tcp", tunnelID),
fmt.Sprintf("%d_mtcp", tunnelID),
}
for _, row := range chainRows {
if row.ChainType == 1 {
_, _ = h.sendNodeCommand(row.NodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
} else if row.ChainType == 2 {
_, _ = h.sendNodeCommand(row.NodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
_, _ = h.sendNodeCommand(row.NodeID, "DeleteService", map[string]interface{}{"services": []string{serviceName}}, false, true)
_, _ = h.sendNodeCommand(row.NodeID, "DeleteService", map[string]interface{}{"services": serviceNames}, false, true)
} else if row.ChainType == 3 {
_, _ = h.sendNodeCommand(row.NodeID, "DeleteService", map[string]interface{}{"services": []string{serviceName}}, false, true)
_, _ = h.sendNodeCommand(row.NodeID, "DeleteService", map[string]interface{}{"services": serviceNames}, false, true)
}
}
}
@@ -816,6 +855,12 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
}
defer func() { tx.Rollback() }()
updateProtocol := "tls"
if len(runtimeState.OutNodes) > 0 && strings.TrimSpace(runtimeState.OutNodes[0].Protocol) != "" {
updateProtocol = strings.TrimSpace(runtimeState.OutNodes[0].Protocol)
} else if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
}
if err := h.repo.UpdateTunnelTx(
tx,
id,
@@ -826,6 +871,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
asInt(req["status"], 1),
inIp,
ipPreference,
updateProtocol,
now,
); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -873,7 +919,11 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
if typeVal == 2 {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil {
h.rollbackTunnelRuntime(createdChains, createdServices, id)
updateProtocol := "tls"
if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
}
h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol)
h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(id)
if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) {
@@ -1696,11 +1746,13 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
return
}
if roleID != 0 {
if roleID != 0 && !h.allowLocalRemoteAddr() {
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络或保留地址"))
response.WriteJSON(w, response.Err(403, err.Error()))
return
}
}
if roleID != 0 {
if speedIDVal, ok := req["speedId"]; ok && speedIDVal != nil {
response.WriteJSON(w, response.Err(-1, "普通用户无法设置限速规则"))
return
@@ -1712,6 +1764,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if roleID != 0 {
if ipSpeedIDVal, ok := req["ipSpeedId"]; ok && ipSpeedIDVal != nil {
response.WriteJSON(w, response.Err(-1, "普通用户无法设置每 IP 限速规则"))
return
}
}
ipSpeedID := asAnyToInt64Ptr(req["ipSpeedId"])
ipSpeedID, err = h.normalizeSpeedLimitReference(ipSpeedID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
port := asInt(req["inPort"], 0)
if port <= 0 {
port = h.pickTunnelPort(tunnelID)
@@ -1749,7 +1813,14 @@ 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, inIp, nullableInt(speedID))
maxConn := asInt(req["maxConn"], 0)
ipMaxConn := asInt(req["ipMaxConn"], 0)
if ipMaxConn < 0 {
ipMaxConn = 0
}
proxyProtocol := asInt(req["proxyProtocol"], 0)
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, ipMaxConn, nullableInt(ipSpeedID), proxyProtocol)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -1820,9 +1891,9 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
if remoteAddr == "" {
remoteAddr = forward.RemoteAddr
}
if actorRole != 0 {
if actorRole != 0 && !h.allowLocalRemoteAddr() {
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络或保留地址"))
response.WriteJSON(w, response.Err(403, err.Error()))
return
}
}
@@ -1848,6 +1919,27 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
} else if _, ok := req["speedId"]; ok {
newSpeedID = sql.NullInt64{Valid: false}
}
rawIPSpeedID, hasIPSpeedID := req["ipSpeedId"]
requestedIPSpeedID := asAnyToInt64Ptr(rawIPSpeedID)
newIPSpeedID := forward.IPSpeedID
if actorRole != 0 {
if hasIPSpeedID && !sameSpeedLimitSelection(forward.IPSpeedID, requestedIPSpeedID) {
response.WriteJSON(w, response.Err(-1, "普通用户无法修改每 IP 限速规则"))
return
}
} else {
ipSpeedID := requestedIPSpeedID
ipSpeedID, err = h.normalizeSpeedLimitReference(ipSpeedID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if ipSpeedID != nil {
newIPSpeedID = sql.NullInt64{Int64: *ipSpeedID, Valid: true}
} else if hasIPSpeedID {
newIPSpeedID = sql.NullInt64{Valid: false}
}
}
port := asInt(req["inPort"], 0)
if port <= 0 {
@@ -1900,7 +1992,14 @@ 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, newSpeedID); err != nil {
maxConn := asInt(req["maxConn"], forward.MaxConn)
ipMaxConn := asInt(req["ipMaxConn"], forward.IPMaxConn)
if ipMaxConn < 0 {
ipMaxConn = 0
}
proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol)
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -3362,6 +3461,11 @@ func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []m
return errors.New("invalid tunnel service context")
}
serviceName := fmt.Sprintf("%d_tls", tunnelID)
if len(serviceData) > 0 {
if name, ok := serviceData[0]["name"].(string); ok && strings.TrimSpace(name) != "" {
serviceName = strings.TrimSpace(name)
}
}
return retryTunnelServiceAddWithCleanup(
func() error {
_, err := h.sendNodeCommand(nodeID, "AddService", serviceData, true, false)
@@ -3375,19 +3479,31 @@ func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []m
)
}
func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tunnelID int64) {
func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tunnelID int64, protocol string) {
if h == nil || tunnelID <= 0 {
return
}
if protocol == "" {
protocol = "tls"
}
seenServices := make(map[int64]struct{})
serviceName := fmt.Sprintf("%d_tls", tunnelID)
serviceNames := []string{
fmt.Sprintf("tunnel_%d", tunnelID),
fmt.Sprintf("%d_tls", tunnelID),
fmt.Sprintf("%d_kcp", tunnelID),
fmt.Sprintf("%d_wss", tunnelID),
fmt.Sprintf("%d_mtls", tunnelID),
fmt.Sprintf("%d_mwss", tunnelID),
fmt.Sprintf("%d_tcp", tunnelID),
fmt.Sprintf("%d_mtcp", tunnelID),
}
for i := len(serviceNodeIDs) - 1; i >= 0; i-- {
nodeID := serviceNodeIDs[i]
if _, ok := seenServices[nodeID]; ok {
continue
}
seenServices[nodeID] = struct{}{}
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": []string{serviceName}}, false, true)
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": serviceNames}, false, true)
}
seenChains := make(map[int64]struct{})
@@ -3457,9 +3573,15 @@ func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRu
connectorMetadata["nodelay"] = true
connectorMetadata["mux.keepaliveInterval"] = "15s"
connectorMetadata["mux.keepaliveTimeout"] = "45s"
connectorMetadata["mux.maxFrameSize"] = 32768
connectorMetadata["mux.maxStreamBuffer"] = 2097152
}
if isKCPTunnelProtocol(protocol) {
connectorMetadata["connectTimeout"] = "30s"
connectorMetadata["mux.keepaliveInterval"] = "15s"
connectorMetadata["mux.keepaliveTimeout"] = "45s"
connectorMetadata["mux.maxFrameSize"] = 32768
connectorMetadata["mux.maxStreamBuffer"] = 2097152
}
if len(connectorMetadata) > 0 {
connector["metadata"] = connectorMetadata
@@ -3505,9 +3627,15 @@ func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode,
handlerMetadata["nodelay"] = true
handlerMetadata["mux.keepaliveInterval"] = "15s"
handlerMetadata["mux.keepaliveTimeout"] = "45s"
handlerMetadata["mux.maxFrameSize"] = 32768
handlerMetadata["mux.maxStreamBuffer"] = 2097152
}
if isKCPTunnelProtocol(protocol) {
handlerMetadata["connectTimeout"] = "30s"
handlerMetadata["mux.keepaliveInterval"] = "15s"
handlerMetadata["mux.keepaliveTimeout"] = "45s"
handlerMetadata["mux.maxFrameSize"] = 32768
handlerMetadata["mux.maxStreamBuffer"] = 2097152
}
if len(handlerMetadata) > 0 {
handlerCfg["metadata"] = handlerMetadata
@@ -3516,7 +3644,7 @@ func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode,
handlerCfg["retries"] = nextHopCandidateCount - 1
}
service := map[string]interface{}{
"name": fmt.Sprintf("%d_%s", tunnelID, protocol),
"name": fmt.Sprintf("tunnel_%d", tunnelID),
"addr": processServerAddress(fmt.Sprintf("%s:%d", defaultString(strings.TrimSpace(chainNode.ConnectIP), node.TCPListenAddr), chainNode.Port)),
"handler": handlerCfg,
"listener": buildTunnelListenerConfig(protocol),
@@ -3621,8 +3749,19 @@ func buildTunnelDialerConfig(protocol string) map[string]interface{} {
}
if isKCPTunnelProtocol(protocol) {
dialer["metadata"] = map[string]interface{}{
"kcp.keepalive": 10,
"kcp.tcp": false,
"kcp.keepalive": 10,
"kcp.tcp": false,
"kcp.mode": "fast3",
"kcp.sndwnd": 4096,
"kcp.rcvwnd": 4096,
"kcp.mtu": 1350,
"kcp.sockbuf": 4194304,
"kcp.smuxbuf": 4194304,
"kcp.streambuf": 2097152,
"kcp.datashard": 10,
"kcp.parityshard": 3,
"kcp.nocomp": true,
"kcp.nc": 1,
}
}
return dialer
@@ -3634,8 +3773,19 @@ func buildTunnelListenerConfig(protocol string) map[string]interface{} {
}
if isKCPTunnelProtocol(protocol) {
listener["metadata"] = map[string]interface{}{
"kcp.keepalive": 10,
"kcp.tcp": false,
"kcp.keepalive": 10,
"kcp.tcp": false,
"kcp.mode": "fast3",
"kcp.sndwnd": 4096,
"kcp.rcvwnd": 4096,
"kcp.mtu": 1350,
"kcp.sockbuf": 4194304,
"kcp.smuxbuf": 4194304,
"kcp.streambuf": 2097152,
"kcp.datashard": 10,
"kcp.parityshard": 3,
"kcp.nocomp": true,
"kcp.nc": 1,
}
}
return listener
@@ -4001,7 +4151,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
h.repo.RollbackForwardFields(
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
oldForward.SpeedID,
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
time.Now().UnixMilli(),
)
@@ -4165,6 +4315,26 @@ func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error {
return nil
}
func (h *Handler) syncUserMaxConnForwards(userID int64) ([]string, error) {
forwards, err := h.listActiveForwardsByUser(userID)
if err != nil {
return nil, err
}
warnings := make([]string, 0)
for i := range forwards {
f := &forwards[i]
if f.MaxConn > 0 {
continue
}
syncWarnings, syncErr := h.syncForwardServicesWithWarnings(f, "UpdateService", true)
warnings = append(warnings, syncWarnings...)
if syncErr != nil {
return warnings, syncErr
}
}
return warnings, nil
}
// cleanupForwardsForUserTunnel deletes all forwarding rules belonging to a
// specific user+tunnel pair. It notifies nodes to remove the runtime services
// first, then deletes the DB records. This is best-effort: individual failures
@@ -11,11 +11,37 @@ var DisableSafeRemoteAddrCheckForTesting = false
// IsSafeRemoteAddr checks if a given address is safe to connect to (prevents SSRF/Open Proxy).
// It resolves domains to IPs to prevent DNS rebinding attacks pointing to internal networks.
// Supports multiple addresses separated by commas or newlines (one per line).
func IsSafeRemoteAddr(addr string) error {
if DisableSafeRemoteAddrCheckForTesting {
return nil
}
for _, part := range splitRemoteParts(addr) {
if err := checkSingleRemoteAddr(part); err != nil {
return err
}
}
return nil
}
// splitRemoteParts splits a multi-address string by commas and newlines.
func splitRemoteParts(addr string) []string {
addr = strings.ReplaceAll(addr, "\n", ",")
addr = strings.ReplaceAll(addr, "\r", ",")
parts := strings.Split(addr, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
part = strings.TrimSpace(part)
if part != "" {
out = append(out, part)
}
}
return out
}
// checkSingleRemoteAddr validates a single address.
func checkSingleRemoteAddr(addr string) error {
host, _, err := net.SplitHostPort(addr)
if err != nil {
if strings.Contains(err.Error(), "missing port in address") {
@@ -27,12 +53,12 @@ func IsSafeRemoteAddr(addr string) error {
ips, err := net.LookupIP(host)
if err != nil {
return fmt.Errorf("could not resolve address: %v", err)
return fmt.Errorf("could not resolve address %q: %v", addr, err)
}
for _, ip := range ips {
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() || ip.IsMulticast() {
return fmt.Errorf("address resolves to internal or reserved IP: %s", ip.String())
if ip.IsLoopback() || ip.IsPrivate() {
return fmt.Errorf("address %q resolves to internal IP: %s", addr, ip.String())
}
}
@@ -6,6 +6,7 @@ import (
"time"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
)
type tunnelTrafficDelta struct {
@@ -21,75 +22,42 @@ func unixMilliBucketMinute(nowMs int64) int64 {
return nowMs - (nowMs % minuteMs)
}
func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowItem, nowMs int64) {
if h == nil || h.repo == nil {
return
}
if nodeID <= 0 || len(items) == 0 {
return
func collectFlowUploadForwardIDs(items []flowItem) []int64 {
ids := make([]int64, 0, len(items))
seen := make(map[int64]struct{}, len(items))
for _, item := range items {
forwardID, _, _, ok := parseFlowServiceIDs(strings.TrimSpace(item.N))
if !ok || forwardID <= 0 {
continue
}
if _, exists := seen[forwardID]; exists {
continue
}
seen[forwardID] = struct{}{}
ids = append(ids, forwardID)
}
return ids
}
func (h *Handler) recordTunnelMetricsFromForwardBatch(nodeID int64, forwardDeltas map[int64]tunnelTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta, nowMs int64) {
if h == nil || h.repo == nil || nodeID <= 0 || len(forwardDeltas) == 0 {
return
}
bucketTs := unixMilliBucketMinute(nowMs)
if bucketTs <= 0 {
return
}
forwardDeltas := make(map[int64]tunnelTrafficDelta)
var skippedParse, skippedZero int
for _, item := range items {
name := strings.TrimSpace(item.N)
if name == "" || name == "web_api" {
continue
}
forwardID, _, _, ok := parseFlowServiceIDs(name)
if !ok {
skippedParse++
continue
}
if item.D == 0 && item.U == 0 {
skippedZero++
continue
}
d := forwardDeltas[forwardID]
d.bytesIn += item.D
d.bytesOut += item.U
forwardDeltas[forwardID] = d
}
if len(forwardDeltas) == 0 {
if len(items) > 0 {
log.Printf("monitoring debug op=tunnel_metric.no_forward_deltas node_id=%d items=%d skipped_parse=%d skipped_zero=%d", nodeID, len(items), skippedParse, skippedZero)
}
return
}
forwardIDs := make([]int64, 0, len(forwardDeltas))
for id := range forwardDeltas {
forwardIDs = append(forwardIDs, id)
}
forwardTunnelMap, err := h.repo.MapForwardIDsToTunnelIDs(forwardIDs)
if err != nil {
log.Printf("monitoring write skipped op=tunnel_metric.map_forward_to_tunnel node_id=%d err=%v", nodeID, err)
return
}
if len(forwardTunnelMap) == 0 {
log.Printf("monitoring debug op=tunnel_metric.no_tunnel_map node_id=%d forward_ids=%v", nodeID, forwardIDs)
return
}
tunnelAgg := make(map[int64]tunnelTrafficDelta)
for forwardID, delta := range forwardDeltas {
tunnelID := forwardTunnelMap[forwardID]
if tunnelID <= 0 {
meta, ok := metas[forwardID]
if !ok || meta.TunnelID <= 0 {
continue
}
a := tunnelAgg[tunnelID]
a.bytesIn += delta.bytesIn
a.bytesOut += delta.bytesOut
tunnelAgg[tunnelID] = a
}
if len(tunnelAgg) == 0 {
return
current := tunnelAgg[meta.TunnelID]
current.bytesIn += delta.bytesIn
current.bytesOut += delta.bytesOut
tunnelAgg[meta.TunnelID] = current
}
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
@@ -98,14 +66,11 @@ func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowIte
continue
}
metrics = append(metrics, &model.TunnelMetric{
TunnelID: tunnelID,
NodeID: nodeID,
Timestamp: bucketTs,
BytesIn: delta.bytesIn,
BytesOut: delta.bytesOut,
Connections: 0,
Errors: 0,
AvgLatencyMs: 0,
TunnelID: tunnelID,
NodeID: nodeID,
Timestamp: bucketTs,
BytesIn: delta.bytesIn,
BytesOut: delta.bytesOut,
})
}
if len(metrics) == 0 {
@@ -114,7 +79,7 @@ func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowIte
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
} else {
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
return
}
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
}
+197 -6
View File
@@ -13,6 +13,13 @@ import (
"go-backend/internal/http/response"
)
// failedForward tracks a forward that failed redeployment, for retry.
type failedForward struct {
id int64
forward *forwardRecord
err error
}
const (
githubRepo = "Sagit-chu/flvx"
githubAPIBase = "https://api.github.com"
@@ -32,6 +39,8 @@ var (
testKeywordPattern = regexp.MustCompile(`(?i)(alpha|beta|rc)`)
)
const nodeOnlineRedeployCooldown = 30 * time.Second
type githubRelease struct {
TagName string `json:"tag_name"`
Name string `json:"name"`
@@ -389,24 +398,123 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
}
func (h *Handler) onNodeOnline(nodeID int64) {
if !h.consumeNodePendingUpgradeRedeploy(nodeID) {
if !h.startNodeOnlineRedeploy(nodeID, time.Now()) {
return
}
h.redeployNodeRuntimeAfterUpgrade(nodeID)
defer h.finishNodeOnlineRedeploy(nodeID)
// Reconcile node runtime on the first reconnect, but suppress rapid flapping
// so websocket churn does not trigger repeated full redeploy storms.
if !h.redeployNodeRuntimeAfterUpgrade(nodeID) {
h.markNodePendingUpgradeRedeploy(nodeID)
}
}
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
func (h *Handler) startNodeOnlineRedeploy(nodeID int64, now time.Time) bool {
if h == nil || nodeID <= 0 {
return false
}
if now.IsZero() {
now = time.Now()
}
h.upgradeMu.Lock()
defer h.upgradeMu.Unlock()
if h.pendingUpgradeRedeploy == nil {
h.pendingUpgradeRedeploy = make(map[int64]struct{})
}
if h.nodeOnlineRedeployAt == nil {
h.nodeOnlineRedeployAt = make(map[int64]time.Time)
}
if h.nodeOnlineRedeployQueued == nil {
h.nodeOnlineRedeployQueued = make(map[int64]struct{})
}
if h.nodeOnlineRedeploying == nil {
h.nodeOnlineRedeploying = make(map[int64]struct{})
}
_, pendingUpgrade := h.pendingUpgradeRedeploy[nodeID]
lastRedeployAt := h.nodeOnlineRedeployAt[nodeID]
_, inFlight := h.nodeOnlineRedeploying[nodeID]
if fireAt, start := nextNodeOnlineRedeployFireAt(lastRedeployAt, now, pendingUpgrade, inFlight); !start {
h.queueNodeOnlineRedeployLocked(nodeID, fireAt)
return false
}
delete(h.pendingUpgradeRedeploy, nodeID)
h.nodeOnlineRedeployAt[nodeID] = now
h.nodeOnlineRedeploying[nodeID] = struct{}{}
return true
}
func nextNodeOnlineRedeployFireAt(lastRedeployAt, now time.Time, pendingUpgrade bool, inFlight bool) (time.Time, bool) {
if now.IsZero() {
now = time.Now()
}
if inFlight {
fireAt := now.Add(nodeOnlineRedeployCooldown)
if !lastRedeployAt.IsZero() {
cooldownAt := lastRedeployAt.Add(nodeOnlineRedeployCooldown)
if cooldownAt.After(now) {
fireAt = cooldownAt
}
}
return fireAt, false
}
if !pendingUpgrade && !lastRedeployAt.IsZero() && now.Sub(lastRedeployAt) < nodeOnlineRedeployCooldown {
return lastRedeployAt.Add(nodeOnlineRedeployCooldown), false
}
return time.Time{}, true
}
func (h *Handler) queueNodeOnlineRedeployLocked(nodeID int64, fireAt time.Time) {
if h == nil || nodeID <= 0 {
return
}
if h.nodeOnlineRedeployQueued == nil {
h.nodeOnlineRedeployQueued = make(map[int64]struct{})
}
if _, queued := h.nodeOnlineRedeployQueued[nodeID]; queued {
return
}
if fireAt.IsZero() {
fireAt = time.Now().Add(nodeOnlineRedeployCooldown)
}
delay := time.Until(fireAt)
if delay < 0 {
delay = 0
}
h.nodeOnlineRedeployQueued[nodeID] = struct{}{}
time.AfterFunc(delay, func() {
h.upgradeMu.Lock()
delete(h.nodeOnlineRedeployQueued, nodeID)
h.upgradeMu.Unlock()
h.onNodeOnline(nodeID)
})
}
func (h *Handler) finishNodeOnlineRedeploy(nodeID int64) {
if h == nil || nodeID <= 0 {
return
}
h.upgradeMu.Lock()
delete(h.nodeOnlineRedeploying, nodeID)
h.upgradeMu.Unlock()
}
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) bool {
tunnelIDs, err := h.repo.ListActiveTunnelIDsByNode(nodeID)
if err != nil {
fmt.Printf("post-upgrade redeploy: list tunnels for node %d failed: %v\n", nodeID, err)
return
return false
}
forwardIDs, err := h.repo.ListActiveForwardIDsByNode(nodeID)
forwardIDs, err := h.repo.ListForwardIDsByNode(nodeID)
if err != nil {
fmt.Printf("post-upgrade redeploy: list forwards for node %d failed: %v\n", nodeID, err)
return
return false
}
// First pass: deploy everything
tunnelFailed := make(map[int64]struct{})
for _, tunnelID := range tunnelIDs {
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
@@ -415,6 +523,9 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
}
}
// Collect forwards that failed independently (not skipped due to tunnel failure)
var failedForwards []failedForward
for _, forwardID := range forwardIDs {
forward, getErr := h.getForwardRecord(forwardID)
if getErr != nil || forward == nil {
@@ -424,7 +535,87 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
continue
}
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
failedForwards = append(failedForwards, failedForward{id: forwardID, forward: forward, err: err})
fmt.Printf("post-upgrade redeploy: forward %d failed on node %d: %v\n", forwardID, nodeID, err)
}
}
// Retry failed items with exponential backoff (max 3 attempts)
return h.retryFailedRedeploys(nodeID, tunnelFailed, failedForwards)
}
// isRetryableError returns true if the error looks transient and worth retrying.
func isRetryableError(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
// Skip non-retryable errors: not-found, already-exists, validation errors
if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") {
return false
}
if strings.Contains(msg, "already exists") || strings.Contains(msg, "已存在") {
return false
}
// Everything else (timeout, connection lost, port in use, etc.) is retryable
return true
}
// retryFailedRedeploys retries failed tunnels and forwards with exponential backoff.
func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]struct{}, failedForwards []failedForward) bool {
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
return true
}
const maxRetries = 3
baseDelay := time.Second
for attempt := 1; attempt <= maxRetries; attempt++ {
delay := baseDelay * time.Duration(1<<uint(attempt-1)) // 1s, 2s, 4s
time.Sleep(delay)
// Retry failed tunnels
for tunnelID := range tunnelFailed {
if err := h.redeployTunnelAndForwards(tunnelID); err == nil {
delete(tunnelFailed, tunnelID)
fmt.Printf("post-upgrade redeploy retry: tunnel %d succeeded on node %d (attempt %d)\n", tunnelID, nodeID, attempt)
} else if !isRetryableError(err) {
delete(tunnelFailed, tunnelID) // Non-retryable, don't retry again
} else {
fmt.Printf("post-upgrade redeploy retry: tunnel %d still failing on node %d (attempt %d): %v\n", tunnelID, nodeID, attempt, err)
}
}
// Retry failed forwards
var stillFailed []failedForward
for _, ff := range failedForwards {
if _, skipped := tunnelFailed[ff.forward.TunnelID]; skipped {
stillFailed = append(stillFailed, ff) // Tunnel still failed, skip forward
continue
}
if err := h.syncForwardServices(ff.forward, "UpdateService", true); err == nil {
fmt.Printf("post-upgrade redeploy retry: forward %d succeeded on node %d (attempt %d)\n", ff.id, nodeID, attempt)
} else if !isRetryableError(err) {
// Non-retryable, drop it
} else {
stillFailed = append(stillFailed, ff)
fmt.Printf("post-upgrade redeploy retry: forward %d still failing on node %d (attempt %d): %v\n", ff.id, nodeID, attempt, err)
}
}
failedForwards = stillFailed
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
fmt.Printf("post-upgrade redeploy retry: all items recovered on node %d\n", nodeID)
return true
}
}
// Final summary
for tunnelID := range tunnelFailed {
fmt.Printf("post-upgrade redeploy: tunnel %d permanently failed on node %d after retries\n", tunnelID, nodeID)
}
for _, ff := range failedForwards {
fmt.Printf("post-upgrade redeploy: forward %d permanently failed on node %d after retries\n", ff.id, nodeID)
}
return false
}
@@ -0,0 +1,111 @@
package handler
import (
"testing"
"time"
)
func TestStartNodeOnlineRedeploySkipsRecentReconnects(t *testing.T) {
h := &Handler{
pendingUpgradeRedeploy: map[int64]struct{}{},
nodeOnlineRedeployAt: map[int64]time.Time{},
nodeOnlineRedeployQueued: map[int64]struct{}{},
nodeOnlineRedeploying: map[int64]struct{}{},
}
now := time.Unix(1_777_176_720, 0)
if !h.startNodeOnlineRedeploy(54, now) {
t.Fatalf("expected first reconnect to redeploy")
}
h.finishNodeOnlineRedeploy(54)
if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
t.Fatalf("expected recent reconnect to skip redeploy")
}
if h.consumeNodePendingUpgradeRedeploy(54) {
t.Fatalf("did not expect pending upgrade marker to be consumed")
}
}
func TestStartNodeOnlineRedeployAllowsPendingUpgradeDuringCooldown(t *testing.T) {
h := &Handler{
pendingUpgradeRedeploy: map[int64]struct{}{},
nodeOnlineRedeployAt: map[int64]time.Time{},
nodeOnlineRedeployQueued: map[int64]struct{}{},
nodeOnlineRedeploying: map[int64]struct{}{},
}
now := time.Unix(1_777_176_720, 0)
if !h.startNodeOnlineRedeploy(54, now) {
t.Fatalf("expected first reconnect to redeploy")
}
h.finishNodeOnlineRedeploy(54)
h.markNodePendingUpgradeRedeploy(54)
if !h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
t.Fatalf("expected pending upgrade reconnect to bypass cooldown")
}
if h.consumeNodePendingUpgradeRedeploy(54) {
t.Fatalf("expected pending upgrade marker to be consumed during redeploy")
}
}
func TestStartNodeOnlineRedeployQueuesCooldownReconnect(t *testing.T) {
h := &Handler{
pendingUpgradeRedeploy: map[int64]struct{}{},
nodeOnlineRedeployAt: map[int64]time.Time{},
nodeOnlineRedeployQueued: map[int64]struct{}{},
nodeOnlineRedeploying: map[int64]struct{}{},
}
now := time.Unix(1_777_176_720, 0)
if !h.startNodeOnlineRedeploy(54, now) {
t.Fatalf("expected first reconnect to redeploy")
}
h.finishNodeOnlineRedeploy(54)
if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
t.Fatalf("expected cooldown reconnect to skip immediate redeploy")
}
if _, queued := h.nodeOnlineRedeployQueued[54]; !queued {
t.Fatalf("expected cooldown reconnect to queue a follow-up redeploy")
}
}
func TestStartNodeOnlineRedeployKeepsPendingUpgradeWhileInFlight(t *testing.T) {
h := &Handler{
pendingUpgradeRedeploy: map[int64]struct{}{},
nodeOnlineRedeployAt: map[int64]time.Time{},
nodeOnlineRedeployQueued: map[int64]struct{}{},
nodeOnlineRedeploying: map[int64]struct{}{},
}
now := time.Unix(1_777_176_720, 0)
if !h.startNodeOnlineRedeploy(54, now) {
t.Fatalf("expected first reconnect to redeploy")
}
h.markNodePendingUpgradeRedeploy(54)
if h.startNodeOnlineRedeploy(54, now.Add(time.Second)) {
t.Fatalf("expected in-flight redeploy to suppress parallel restart")
}
if !h.consumeNodePendingUpgradeRedeploy(54) {
t.Fatalf("expected pending upgrade marker to remain for the next retry")
}
h.finishNodeOnlineRedeploy(54)
}
func TestNextNodeOnlineRedeployFireAtDefersExpiredInFlightReconnect(t *testing.T) {
now := time.Unix(1_777_176_720, 0)
last := now.Add(-nodeOnlineRedeployCooldown - 5*time.Second)
fireAt, start := nextNodeOnlineRedeployFireAt(last, now, false, true)
if start {
t.Fatalf("expected in-flight reconnect to queue instead of starting immediately")
}
want := now.Add(nodeOnlineRedeployCooldown)
if !fireAt.Equal(want) {
t.Fatalf("expected queued reconnect at %s, got %s", want, fireAt)
}
}
+51 -38
View File
@@ -23,26 +23,31 @@ type User struct {
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
Status int `gorm:"not null"`
MaxConn int `gorm:"column:max_conn;not null;default:0"`
}
func (User) TableName() string { return "user" }
// Forward maps to the "forward" table.
type Forward struct {
ID int64 `gorm:"primaryKey;autoIncrement"`
UserID int64 `gorm:"column:user_id;not null"`
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
Name string `gorm:"type:varchar(100);not null"`
TunnelID int64 `gorm:"column:tunnel_id;not null"`
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
InFlow int64 `gorm:"not null;default:0"`
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime int64 `gorm:"column:updated_time;not null"`
Status int `gorm:"not null"`
Inx int `gorm:"not null;default:0"`
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
ID int64 `gorm:"primaryKey;autoIncrement"`
UserID int64 `gorm:"column:user_id;not null"`
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
Name string `gorm:"type:varchar(100);not null"`
TunnelID int64 `gorm:"column:tunnel_id;not null"`
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
InFlow int64 `gorm:"not null;default:0"`
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime int64 `gorm:"column:updated_time;not null"`
Status int `gorm:"not null"`
Inx int `gorm:"not null;default:0"`
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
MaxConn int `gorm:"column:max_conn;not null;default:0"`
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"`
}
func (Forward) TableName() string { return "forward" }
@@ -425,21 +430,24 @@ type ChainTunnelBackup struct {
}
type ForwardBackup struct {
ID int64 `json:"id"`
UserID int64 `json:"userId"`
UserName string `json:"userName"`
Name string `json:"name"`
TunnelID int64 `json:"tunnelId"`
RemoteAddr string `json:"remoteAddr"`
Strategy string `json:"strategy"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
Inx int `json:"inx"`
SpeedID *int64 `json:"speedId,omitempty"`
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
ID int64 `json:"id"`
UserID int64 `json:"userId"`
UserName string `json:"userName"`
Name string `json:"name"`
TunnelID int64 `json:"tunnelId"`
RemoteAddr string `json:"remoteAddr"`
Strategy string `json:"strategy"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
Inx int `json:"inx"`
SpeedID *int64 `json:"speedId,omitempty"`
IPMaxConn int `json:"ipMaxConn,omitempty"`
IPSpeedID *int64 `json:"ipSpeedId,omitempty"`
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
ProxyProtocol int `json:"proxyProtocol"`
}
type ForwardPortBackup struct {
@@ -528,15 +536,19 @@ type ImportResult struct {
// ForwardRecord is a minimal forward view used by control plane and flow policy.
type ForwardRecord struct {
ID int64
UserID int64
UserName string
Name string
TunnelID int64
RemoteAddr string
Strategy string
Status int
SpeedID sql.NullInt64
ID int64
UserID int64
UserName string
Name string
TunnelID int64
RemoteAddr string
Strategy string
Status int
SpeedID sql.NullInt64
MaxConn int
IPMaxConn int
IPSpeedID sql.NullInt64
ProxyProtocol int
}
// TunnelRecord is a minimal tunnel view used by control plane.
@@ -546,6 +558,7 @@ type TunnelRecord struct {
Status int
Flow int64
TrafficRatio float64
Protocol string
}
type UserQuotaView struct {
+183 -36
View File
@@ -22,6 +22,13 @@ import (
"go-backend/internal/store/model"
)
const (
defaultPostgresMaxOpenConns = 32
defaultPostgresMaxIdleConns = 8
defaultPostgresConnMaxIdle = 5 * time.Minute
defaultPostgresConnMaxLife = 30 * time.Minute
)
// ─── Type aliases for backward compatibility ─────────────────────────
// Handlers still reference repo.User, repo.BackupData, etc.
@@ -61,6 +68,14 @@ type Repository struct {
db *gorm.DB
}
type FlowUploadCounterDelta struct {
ForwardID int64
UserID int64
UserTunnelID int64
InFlow int64
OutFlow int64
}
func (r *Repository) DB() *gorm.DB {
if r == nil {
return nil
@@ -68,6 +83,79 @@ func (r *Repository) DB() *gorm.DB {
return r.db
}
func sortedFlowUploadTargetIDs(totals map[int64][2]int64) []int64 {
ids := make([]int64, 0, len(totals))
for id := range totals {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids
}
func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if len(deltas) == 0 {
return nil
}
forwardTotals := make(map[int64][2]int64, len(deltas))
userTotals := make(map[int64][2]int64, len(deltas))
userTunnelTotals := make(map[int64][2]int64, len(deltas))
for _, delta := range deltas {
if delta.ForwardID > 0 {
current := forwardTotals[delta.ForwardID]
current[0] += delta.InFlow
current[1] += delta.OutFlow
forwardTotals[delta.ForwardID] = current
}
if delta.UserID > 0 {
current := userTotals[delta.UserID]
current[0] += delta.InFlow
current[1] += delta.OutFlow
userTotals[delta.UserID] = current
}
if delta.UserTunnelID > 0 {
current := userTunnelTotals[delta.UserTunnelID]
current[0] += delta.InFlow
current[1] += delta.OutFlow
userTunnelTotals[delta.UserTunnelID] = current
}
}
return r.db.Transaction(func(tx *gorm.DB) error {
for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) {
total := forwardTotals[forwardID]
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{
"in_flow": gorm.Expr("in_flow + ?", total[0]),
"out_flow": gorm.Expr("out_flow + ?", total[1]),
}).Error; err != nil {
return err
}
}
for _, userID := range sortedFlowUploadTargetIDs(userTotals) {
total := userTotals[userID]
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{
"in_flow": gorm.Expr("in_flow + ?", total[0]),
"out_flow": gorm.Expr("out_flow + ?", total[1]),
}).Error; err != nil {
return err
}
}
for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) {
total := userTunnelTotals[userTunnelID]
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{
"in_flow": gorm.Expr("in_flow + ?", total[0]),
"out_flow": gorm.Expr("out_flow + ?", total[1]),
}).Error; err != nil {
return err
}
}
return nil
})
}
// ─── Open / Close ────────────────────────────────────────────────────
func Open(path string) (*Repository, error) {
@@ -129,6 +217,7 @@ func OpenPostgres(dsn string) (*Repository, error) {
if err != nil {
return nil, err
}
configurePostgresPool(sqlDB)
if err := sqlDB.Ping(); err != nil {
_ = sqlDB.Close()
return nil, err
@@ -154,6 +243,16 @@ func OpenPostgres(dsn string) (*Repository, error) {
return &Repository{db: db}, nil
}
func configurePostgresPool(sqlDB *sql.DB) {
if sqlDB == nil {
return
}
sqlDB.SetMaxOpenConns(defaultPostgresMaxOpenConns)
sqlDB.SetMaxIdleConns(defaultPostgresMaxIdleConns)
sqlDB.SetConnMaxIdleTime(defaultPostgresConnMaxIdle)
sqlDB.SetConnMaxLifetime(defaultPostgresConnMaxLife)
}
func (r *Repository) Close() error {
if r == nil || r.db == nil {
return nil
@@ -295,6 +394,17 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
}
}
if m.HasTable(&model.Forward{}) {
for _, field := range []string{"ProxyProtocol"} {
if m.HasColumn(&model.Forward{}, field) {
continue
}
if err := m.AddColumn(&model.Forward{}, field); err != nil {
return fmt.Errorf("add forward.%s: %w", field, err)
}
}
}
return nil
}
@@ -679,11 +789,11 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
"version": nullableString(n.Version),
"http": n.HTTP, "tls": n.TLS, "socks": n.Socks,
"status": n.Status, "isRemote": n.IsRemote,
"remoteUrl": nullableString(n.RemoteURL),
"remoteToken": nullableString(n.RemoteToken),
"remoteConfig": nullableString(n.RemoteConfig),
"expiryReminderDismissed": n.ExpiryReminderDismissed,
"interfaceName": nullableString(n.InterfaceName),
"remoteUrl": nullableString(n.RemoteURL),
"remoteToken": nullableString(n.RemoteToken),
"remoteConfig": nullableString(n.RemoteConfig),
"expiryReminderDismissed": n.ExpiryReminderDismissed,
"interfaceName": nullableString(n.InterfaceName),
})
}
return items, nil
@@ -714,6 +824,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
"updatedTime": nullableInt64(u.UpdatedTime),
"inFlow": u.InFlow, "outFlow": u.OutFlow,
"maxConn": u.MaxConn,
}
if quota := quotaMap[u.ID]; quota != nil {
item["dailyQuotaGB"] = quota.DailyLimitGB
@@ -754,27 +865,33 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
}
type fwdRow struct {
ID int64
UserID int64
UserName string
Name string
TunnelID int64
TunnelName string
TrafficRatio float64
RemoteAddr string
Strategy string
InFlow int64
OutFlow int64
CreatedTime int64
Status int
Inx int
SpeedID sql.NullInt64
ID int64
UserID int64
UserName string
Name string
TunnelID int64
TunnelName string
TrafficRatio float64
RemoteAddr string
Strategy string
InFlow int64
OutFlow int64
CreatedTime int64
Status int
Inx int
SpeedID sql.NullInt64
MaxConn int
IPMaxConn int
IPSpeedID sql.NullInt64
IPSpeedLimitName string
ProxyProtocol int
}
var rows []fwdRow
err := r.db.Model(&model.Forward{}).
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id").
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.ip_max_conn, forward.ip_speed_id, COALESCE(ip_speed_limit.name, '') AS ip_speed_limit_name, forward.proxy_protocol").
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
Joins("LEFT JOIN speed_limit AS ip_speed_limit ON ip_speed_limit.id = forward.ip_speed_id").
Order("forward.inx ASC, forward.id ASC").
Find(&rows).Error
if err != nil {
@@ -795,10 +912,19 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
"inFlow": row.InFlow, "outFlow": row.OutFlow,
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
"maxConn": row.MaxConn,
"ipMaxConn": row.IPMaxConn,
"proxyProtocol": row.ProxyProtocol,
}
if row.SpeedID.Valid {
item["speedId"] = row.SpeedID.Int64
}
if row.IPSpeedID.Valid {
item["ipSpeedId"] = row.IPSpeedID.Int64
}
if strings.TrimSpace(row.IPSpeedLimitName) != "" {
item["ipSpeedLimitName"] = row.IPSpeedLimitName
}
items = append(items, item)
}
return items, nil
@@ -1987,6 +2113,16 @@ func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
IPMaxConn: f.IPMaxConn,
ProxyProtocol: f.ProxyProtocol,
}
if f.SpeedID.Valid {
v := f.SpeedID.Int64
b.SpeedID = &v
}
if f.IPSpeedID.Valid {
v := f.IPSpeedID.Int64
b.IPSpeedID = &v
}
ports, err := r.exportForwardPorts(f.ID)
if err != nil {
@@ -2367,29 +2503,40 @@ func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, e
return count, nil
}
func nullableBackupInt64(v *int64) int64 {
if v == nil {
return 0
}
return *v
}
func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int, error) {
count := 0
for _, f := range forwards {
item := model.Forward{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
InFlow: f.InFlow,
OutFlow: f.OutFlow,
CreatedTime: f.CreatedTime,
UpdatedTime: now,
Status: f.Status,
Inx: f.Inx,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
InFlow: f.InFlow,
OutFlow: f.OutFlow,
CreatedTime: f.CreatedTime,
UpdatedTime: now,
Status: f.Status,
Inx: f.Inx,
SpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.SpeedID), Valid: f.SpeedID != nil && *f.SpeedID > 0},
IPMaxConn: f.IPMaxConn,
IPSpeedID: sql.NullInt64{Int64: nullableBackupInt64(f.IPSpeedID), Valid: f.IPSpeedID != nil && *f.IPSpeedID > 0},
ProxyProtocol: f.ProxyProtocol,
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
"user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy",
"in_flow", "out_flow", "updated_time", "status", "inx",
"in_flow", "out_flow", "updated_time", "status", "inx", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol",
}),
}).Create(&item).Error
if err != nil {
@@ -3337,7 +3484,7 @@ func (r *Repository) GetNodeMetrics(nodeID int64, startMs, endMs int64) ([]model
rangeMs := endMs - startMs
const maxRawRangeMs = int64(60 * 60 * 1000) // 1 hour — return raw data for short ranges
const targetPoints = 500 // target number of chart points for downsampled data
const targetPoints = 500 // target number of chart points for downsampled data
// For short ranges, return raw data (full resolution).
if rangeMs <= maxRawRangeMs {
@@ -45,15 +45,19 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
rows := make([]model.ForwardRecord, 0, len(forwards))
for _, f := range forwards {
rows = append(rows, model.ForwardRecord{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
for i := range rows {
@@ -64,7 +68,6 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
return rows, nil
}
func (r *Repository) ListActiveTunnelIDsByNode(nodeID int64) ([]int64, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
@@ -99,6 +102,22 @@ func (r *Repository) ListActiveForwardIDsByNode(nodeID int64) ([]int64, error) {
return ids, nil
}
func (r *Repository) ListForwardIDsByNode(nodeID int64) ([]int64, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ids []int64
err := r.db.Model(&model.ForwardPort{}).
Where("forward_port.node_id = ?", nodeID).
Select("DISTINCT forward_port.forward_id").
Order("forward_port.forward_id ASC").
Pluck("forward_port.forward_id", &ids).Error
if err != nil {
return nil, err
}
return ids, nil
}
func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecord, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
@@ -126,7 +145,6 @@ func (r *Repository) ListForwardPortsTx(tx *gorm.DB, forwardID int64) ([]model.F
return rows, nil
}
func (r *Repository) HasOtherForwardOnNodePort(nodeID int64, port int, currentForwardID int64) (bool, error) {
if r == nil || r.db == nil {
return false, errors.New("repository not initialized")
@@ -153,7 +171,6 @@ func (r *Repository) HasOtherForwardOnNodePortTx(tx *gorm.DB, nodeID int64, port
return count > 0, nil
}
func (r *Repository) GetTunnelOutProtocol(tunnelID int64) (string, error) {
if r == nil || r.db == nil {
return "", errors.New("repository not initialized")
+137 -36
View File
@@ -9,6 +9,90 @@ import (
"go-backend/internal/store/model"
)
type FlowUploadForwardMeta struct {
ForwardID int64
TunnelID int64
TrafficRatio float64
TunnelFlow int64
}
const flowUploadForwardMetaChunkSize = 500
func chunkFlowUploadForwardIDs(ids []int64) [][]int64 {
if len(ids) == 0 {
return nil
}
chunks := make([][]int64, 0, (len(ids)+flowUploadForwardMetaChunkSize-1)/flowUploadForwardMetaChunkSize)
for start := 0; start < len(ids); start += flowUploadForwardMetaChunkSize {
end := start + flowUploadForwardMetaChunkSize
if end > len(ids) {
end = len(ids)
}
chunks = append(chunks, ids[start:end])
}
return chunks
}
func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]FlowUploadForwardMeta, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if len(forwardIDs) == 0 {
return map[int64]FlowUploadForwardMeta{}, nil
}
ids := make([]int64, 0, len(forwardIDs))
seen := make(map[int64]struct{}, len(forwardIDs))
for _, id := range forwardIDs {
if id <= 0 {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
ids = append(ids, id)
}
if len(ids) == 0 {
return map[int64]FlowUploadForwardMeta{}, nil
}
type row struct {
ForwardID int64 `gorm:"column:forward_id"`
TunnelID int64 `gorm:"column:tunnel_id"`
TrafficRatio float64 `gorm:"column:traffic_ratio"`
TunnelFlow int64 `gorm:"column:tunnel_flow"`
}
out := make(map[int64]FlowUploadForwardMeta, len(ids))
for _, chunk := range chunkFlowUploadForwardIDs(ids) {
var rows []row
err := r.db.Table("forward AS f").
Select("f.id AS forward_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
Joins("LEFT JOIN tunnel t ON t.id = f.tunnel_id").
Where("f.id IN ?", chunk).
Scan(&rows).Error
if err != nil {
return nil, err
}
for _, row := range rows {
if row.TunnelFlow <= 0 {
row.TunnelFlow = 1
}
if row.TrafficRatio <= 0 {
row.TrafficRatio = 1
}
out[row.ForwardID] = FlowUploadForwardMeta{
ForwardID: row.ForwardID,
TunnelID: row.TunnelID,
TrafficRatio: row.TrafficRatio,
TunnelFlow: row.TunnelFlow,
}
}
}
return out, nil
}
func (r *Repository) UpdateForwardStatus(forwardID int64, status int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
@@ -30,15 +114,19 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
rows := make([]model.ForwardRecord, 0, len(forwards))
for _, f := range forwards {
rows = append(rows, model.ForwardRecord{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
for i := range rows {
@@ -61,15 +149,19 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
rows := make([]model.ForwardRecord, 0, len(forwards))
for _, f := range forwards {
rows = append(rows, model.ForwardRecord{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
for i := range rows {
@@ -92,15 +184,19 @@ func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]mode
rows := make([]model.ForwardRecord, 0, len(forwards))
for _, f := range forwards {
rows = append(rows, model.ForwardRecord{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
})
}
for i := range rows {
@@ -124,15 +220,19 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
return nil, err
}
fr := model.ForwardRecord{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
}
if strings.TrimSpace(fr.Strategy) == "" {
fr.Strategy = "fifo"
@@ -158,6 +258,7 @@ func (r *Repository) GetTunnelRecord(tunnelID int64) (*model.TunnelRecord, error
Status: t.Status,
Flow: t.Flow,
TrafficRatio: t.TrafficRatio,
Protocol: t.Protocol,
}
if tr.Flow <= 0 {
tr.Flow = 1
@@ -0,0 +1,142 @@
package repo
import (
"path/filepath"
"reflect"
"testing"
"time"
)
func TestChunkFlowUploadForwardIDs(t *testing.T) {
ids := make([]int64, 0, 1001)
for i := int64(1); i <= 1001; i++ {
ids = append(ids, i)
}
chunks := chunkFlowUploadForwardIDs(ids)
if len(chunks) != 3 {
t.Fatalf("expected 3 chunks, got %d", len(chunks))
}
if len(chunks[0]) != 500 || len(chunks[1]) != 500 || len(chunks[2]) != 1 {
t.Fatalf("unexpected chunk sizes: %d, %d, %d", len(chunks[0]), len(chunks[1]), len(chunks[2]))
}
if chunks[0][0] != 1 || chunks[1][0] != 501 || chunks[2][0] != 1001 {
t.Fatalf("unexpected chunk boundaries: %#v %#v %#v", chunks[0][:1], chunks[1][:1], chunks[2][:1])
}
}
func TestSortedFlowUploadTargetIDs(t *testing.T) {
totals := map[int64][2]int64{
9: {1, 1},
2: {1, 1},
7: {1, 1},
}
got := sortedFlowUploadTargetIDs(totals)
want := []int64{2, 7, 9}
if !reflect.DeepEqual(got, want) {
t.Fatalf("expected sorted ids %v, got %v", want, got)
}
}
func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "flow-batch.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := r.DB().Exec(`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
metas, err := r.GetFlowUploadForwardMetas([]int64{20, 99})
if err != nil {
t.Fatalf("get metas: %v", err)
}
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
t.Fatalf("unexpected meta for forward 20: %#v", metas[20])
}
if _, ok := metas[99]; ok {
t.Fatalf("did not expect meta for missing forward 99")
}
err = r.ApplyFlowUploadDeltasBatch([]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}})
if err != nil {
t.Fatalf("apply flow batch: %v", err)
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
t.Fatalf("expected forward in_flow=480, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
t.Fatalf("expected user out_flow=660, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
}
}
func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(25, 2, 'u2', 'f25', 99, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
metas, err := r.GetFlowUploadForwardMetas([]int64{25})
if err != nil {
t.Fatalf("get metas: %v", err)
}
meta, ok := metas[25]
if !ok {
t.Fatalf("expected metadata for forward with missing tunnel row")
}
if meta.ForwardID != 25 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
t.Fatalf("unexpected fallback meta: %#v", meta)
}
}
func TestAddUserQuotaUsageBatchReturnsNormalizedViews(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "quota-batch.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
views, err := r.AddUserQuotaUsageBatch(map[int64]int64{2: 1140}, now)
if err != nil {
t.Fatalf("batch quota update: %v", err)
}
if views[2] == nil || views[2].DailyUsedBytes != 1140 || views[2].MonthlyUsedBytes != 1140 {
t.Fatalf("unexpected quota view: %#v", views[2])
}
}
func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...interface{}) int64 {
t.Helper()
var value int64
if err := r.DB().Raw(query, args...).Row().Scan(&value); err != nil {
t.Fatalf("query %q failed: %v", query, err)
}
return value
}
@@ -0,0 +1,259 @@
package repo
import (
"database/sql"
"testing"
"time"
"go-backend/internal/store/model"
)
func TestGetForwardRecordIncludesProxyProtocol(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Create(&model.Forward{
UserID: 1,
UserName: "admin",
Name: "proxy-forward",
TunnelID: 1,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
ProxyProtocol: 2,
}).Error; err != nil {
t.Fatalf("create forward: %v", err)
}
forwardID := mustRepoLastInsertID(t, r)
record, err := r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord: %v", err)
}
if record == nil {
t.Fatalf("expected forward record")
}
if record.ProxyProtocol != 2 {
t.Fatalf("expected proxyProtocol 2, got %d", record.ProxyProtocol)
}
if record.MaxConn != 0 {
t.Fatalf("expected default maxConn 0, got %d", record.MaxConn)
}
}
func TestListForwardsByTunnelIncludesProxyProtocol(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Create(&model.Forward{
UserID: 1,
UserName: "admin",
Name: "proxy-forward",
TunnelID: 7,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
ProxyProtocol: 2,
}).Error; err != nil {
t.Fatalf("create forward: %v", err)
}
records, err := r.ListForwardsByTunnel(7)
if err != nil {
t.Fatalf("ListForwardsByTunnel: %v", err)
}
if len(records) != 1 {
t.Fatalf("expected 1 forward record, got %d", len(records))
}
if records[0].ProxyProtocol != 2 {
t.Fatalf("expected proxyProtocol 2, got %d", records[0].ProxyProtocol)
}
}
func TestListForwardsByTunnelIncludesMaxConn(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Create(&model.Forward{
UserID: 1,
UserName: "admin",
Name: "max-conn-forward",
TunnelID: 8,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
MaxConn: 42,
}).Error; err != nil {
t.Fatalf("create forward: %v", err)
}
records, err := r.ListForwardsByTunnel(8)
if err != nil {
t.Fatalf("ListForwardsByTunnel: %v", err)
}
if len(records) != 1 {
t.Fatalf("expected 1 forward record, got %d", len(records))
}
if records[0].MaxConn != 42 {
t.Fatalf("expected maxConn 42, got %d", records[0].MaxConn)
}
}
func TestListActiveForwardsByUserTunnelIncludesMaxConn(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.DB().Create(&model.Forward{
UserID: 2,
UserName: "user",
Name: "active-max-conn-forward",
TunnelID: 9,
RemoteAddr: "1.1.1.1:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
MaxConn: 55,
}).Error; err != nil {
t.Fatalf("create forward: %v", err)
}
records, err := r.ListActiveForwardsByUserTunnel(2, 9)
if err != nil {
t.Fatalf("ListActiveForwardsByUserTunnel: %v", err)
}
if len(records) != 1 {
t.Fatalf("expected 1 forward record, got %d", len(records))
}
if records[0].MaxConn != 55 {
t.Fatalf("expected maxConn 55, got %d", records[0].MaxConn)
}
}
func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
forwardID, err := r.CreateForwardTx(1, "admin", "per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, []int64{3}, 24000, "", nil, 0, 5, int64(21), 0)
if err != nil {
t.Fatalf("CreateForwardTx: %v", err)
}
record, err := r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord after create: %v", err)
}
if record.IPMaxConn != 5 {
t.Fatalf("expected created ipMaxConn 5, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
t.Fatalf("expected created ipSpeedId 21, got %+v", record.IPSpeedID)
}
if err := r.UpdateForward(forwardID, "per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 9, int64(22), 0); err != nil {
t.Fatalf("UpdateForward: %v", err)
}
record, err = r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord after update: %v", err)
}
if record.IPMaxConn != 9 {
t.Fatalf("expected updated ipMaxConn 9, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 22 {
t.Fatalf("expected updated ipSpeedId 22, got %+v", record.IPSpeedID)
}
if err := r.DB().Create(&model.Forward{
UserID: 4,
UserName: "user",
Name: "listed-per-ip-forward",
TunnelID: 8,
RemoteAddr: "3.3.3.3:443",
Strategy: "fifo",
CreatedTime: now,
UpdatedTime: now,
Status: 1,
IPMaxConn: 11,
IPSpeedID: sql.NullInt64{Int64: 33, Valid: true},
}).Error; err != nil {
t.Fatalf("create listed forward: %v", err)
}
records, err := r.ListForwardsByTunnel(8)
if err != nil {
t.Fatalf("ListForwardsByTunnel: %v", err)
}
if len(records) != 1 {
t.Fatalf("expected 1 listed record, got %d", len(records))
}
if records[0].IPMaxConn != 11 || !records[0].IPSpeedID.Valid || records[0].IPSpeedID.Int64 != 33 {
t.Fatalf("expected listed per-IP limits 11/33, got ipMaxConn=%d ipSpeedId=%+v", records[0].IPMaxConn, records[0].IPSpeedID)
}
}
func TestRollbackForwardFieldsRestoresPerIPLimits(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
forwardID, err := r.CreateForwardTx(1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", now, 1, nil, 0, "", nil, 7, 5, int64(21), 2)
if err != nil {
t.Fatalf("CreateForwardTx: %v", err)
}
if err := r.UpdateForward(forwardID, "rollback-per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0); err != nil {
t.Fatalf("UpdateForward: %v", err)
}
r.RollbackForwardFields(forwardID, 1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", 1, nil, 7, 5, int64(21), 2, now+2)
record, err := r.GetForwardRecord(forwardID)
if err != nil {
t.Fatalf("GetForwardRecord: %v", err)
}
if record.IPMaxConn != 5 {
t.Fatalf("expected rollback ipMaxConn 5, got %d", record.IPMaxConn)
}
if !record.IPSpeedID.Valid || record.IPSpeedID.Int64 != 21 {
t.Fatalf("expected rollback ipSpeedId 21, got %+v", record.IPSpeedID)
}
}
func mustRepoLastInsertID(t *testing.T, r *Repository) int64 {
t.Helper()
var id int64
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
t.Fatalf("last_insert_rowid: %v", err)
}
if id <= 0 {
t.Fatalf("invalid last_insert_rowid %d", id)
}
return id
}
@@ -37,7 +37,7 @@ func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool
return cnt > 0, err
}
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status int, now int64) (int64, error) {
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
@@ -51,6 +51,7 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
OutFlow: 0,
FlowResetTime: flowResetTime,
Num: num,
MaxConn: maxConn,
CreatedTime: now,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: status,
@@ -73,7 +74,7 @@ func (r *Repository) GetUserRoleID(userID int64) (int, error) {
return user.RoleID, nil
}
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status int, now int64) error {
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -87,11 +88,12 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
"exp_time": expTime,
"flow_reset_time": flowResetTime,
"status": status,
"max_conn": maxConn,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
}).Error
}
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status int, now int64) error {
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -104,6 +106,7 @@ func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow i
"exp_time": expTime,
"flow_reset_time": flowResetTime,
"status": status,
"max_conn": maxConn,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
}).Error
}
@@ -394,7 +397,7 @@ func (r *Repository) UpdateTunnelOrder(tunnelID int64, inx int, now int64) {
Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error
}
func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, typeVal int, flow int64, trafficRatio float64, status int, inIP, ipPreference string, now int64) error {
func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, typeVal int, flow int64, trafficRatio float64, status int, inIP, ipPreference string, protocol string, now int64) error {
if tx == nil {
return errors.New("database unavailable")
}
@@ -408,6 +411,7 @@ func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, ty
"status": status,
"in_ip": nullStringFromInterface(inIP),
"ip_preference": ipPreference,
"protocol": protocol,
"updated_time": now,
}).Error
}
@@ -691,19 +695,23 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
return p
}
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}) error {
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.Forward{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"speed_id": nullInt64FromInterface(speedID),
"updated_time": now,
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol,
"updated_time": now,
}).Error
}
@@ -777,22 +785,26 @@ func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int,
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
}
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, now int64) {
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, now int64) {
if r == nil || r.db == nil {
return
}
_ = r.db.Model(&model.Forward{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"user_id": userID,
"user_name": userName,
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"status": status,
"speed_id": nullInt64FromInterface(speedID),
"updated_time": now,
"user_id": userID,
"user_name": userName,
"name": name,
"tunnel_id": tunnelID,
"remote_addr": remoteAddr,
"strategy": strategy,
"status": status,
"speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn,
"ip_max_conn": ipMaxConn,
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
"proxy_protocol": proxyProtocol,
"updated_time": now,
}).Error
}
@@ -1252,26 +1264,30 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
return ut.ID, true, nil
}
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}) (int64, error) {
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
var forwardID int64
err := r.db.Transaction(func(tx *gorm.DB) error {
fwd := model.Forward{
UserID: userID,
UserName: userName,
Name: name,
TunnelID: tunnelID,
RemoteAddr: remoteAddr,
Strategy: strategy,
InFlow: 0,
OutFlow: 0,
CreatedTime: now,
UpdatedTime: now,
Status: 1,
Inx: inx,
SpeedID: nullInt64FromInterface(speedID),
UserID: userID,
UserName: userName,
Name: name,
TunnelID: tunnelID,
RemoteAddr: remoteAddr,
Strategy: strategy,
InFlow: 0,
OutFlow: 0,
CreatedTime: now,
UpdatedTime: now,
Status: 1,
Inx: inx,
MaxConn: maxConn,
SpeedID: nullInt64FromInterface(speedID),
IPMaxConn: ipMaxConn,
IPSpeedID: nullInt64FromInterface(ipSpeedID),
ProxyProtocol: proxyProtocol,
}
if err := tx.Create(&fwd).Error; err != nil {
return err
@@ -0,0 +1,32 @@
package repo
import (
"testing"
gsqlite "github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
func TestConfigurePostgresPoolSetsMaxOpenConnections(t *testing.T) {
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db handle: %v", err)
}
t.Cleanup(func() {
_ = sqlDB.Close()
})
configurePostgresPool(sqlDB)
if got := sqlDB.Stats().MaxOpenConnections; got != defaultPostgresMaxOpenConns {
t.Fatalf("expected max open conns %d, got %d", defaultPostgresMaxOpenConns, got)
}
}
@@ -3,6 +3,7 @@ package repo
import (
"errors"
"fmt"
"sort"
"strconv"
"strings"
"time"
@@ -255,6 +256,54 @@ func (r *Repository) AddUserQuotaUsage(userID int64, usedBytes int64, now time.T
return normalizeUserQuotaView(result, now), nil
}
func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if len(usages) == 0 {
return map[int64]*model.UserQuotaView{}, nil
}
result := make(map[int64]*model.UserQuotaView, len(usages))
err := r.db.Transaction(func(tx *gorm.DB) error {
userIDs := make([]int64, 0, len(usages))
for userID := range usages {
if userID > 0 {
userIDs = append(userIDs, userID)
}
}
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
for _, userID := range userIDs {
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
if err != nil {
return err
}
applyUserQuotaWindowRoll(q, now)
if usages[userID] > 0 {
q.DailyUsedBytes += usages[userID]
q.MonthlyUsedBytes += usages[userID]
}
q.UpdatedTime = now.UnixMilli()
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
"daily_used_bytes": q.DailyUsedBytes,
"monthly_used_bytes": q.MonthlyUsedBytes,
"day_key": q.DayKey,
"month_key": q.MonthKey,
"updated_time": q.UpdatedTime,
}).Error; err != nil {
return err
}
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
}
return nil
})
if err != nil {
return nil, err
}
return result, nil
}
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
+163
View File
@@ -0,0 +1,163 @@
import re
with open('internal/http/handler/control_plane.go', 'r') as f:
content = f.read()
# 1. Update ensureLimiterOnNode and add ensureConnLimiterOnNode
ensure_conn_limiter = """
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error {
\tlimitStr := fmt.Sprintf("$ %d", maxConn)
\t
\tpayload := map[string]interface{}{
\t\t"name": limiterName,
\t\t"limits": []string{limitStr},
\t}
\t
\tif _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
\t\tif !isAlreadyExistsMessage(err.Error()) {
\t\t\treturn fmt.Errorf("连接限制器下发失败: %w", err)
\t\t}
\t\tupdatePayload := map[string]interface{}{
\t\t\t"limiter": limiterName,
\t\t\t"data": payload,
\t\t}
\t\tif _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
\t\t\treturn fmt.Errorf("连接限制器更新失败: %w", updateErr)
\t\t}
\t}
\treturn nil
}
"""
content = content.replace('func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {\n\tif err := h.upsertLimiterOnNode(nodeID, limiterID, speed); err != nil {\n\t\treturn fmt.Errorf("限速规则下发失败: %w", err)\n\t}\n\n\treturn nil\n}',
'func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {\n\tif err := h.upsertLimiterOnNode(nodeID, limiterID, speed); err != nil {\n\t\treturn fmt.Errorf("限速规则下发失败: %w", err)\n\t}\n\n\treturn nil\n}\n' + ensure_conn_limiter)
# 2. Update buildForwardServiceConfigs declaration
content = content.replace(
'func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64) []map[string]interface{} {',
'func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, cLimiterName string) []map[string]interface{} {'
)
# 3. Inject climiter into generated service
service_map_end = """ }
if protocol == "udp" {"""
service_map_end_new = """ }
if cLimiterName != "" {
service["climiter"] = cLimiterName
}
if protocol == "udp" {"""
content = content.replace(service_map_end, service_map_end_new)
# 4. Update syncForwardServicesWithWarnings
# Find user tunnel resolution
resolution = """ serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
for _, fp := range ports {"""
resolution_new = """ serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
user, err := h.repo.GetUserByID(forward.UserID)
if err != nil {
return nil, err
}
var cLimiterName string
var maxConnToSet int
if forward.MaxConn > 0 {
maxConnToSet = forward.MaxConn
cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID)
} else if user != nil && user.MaxConn > 0 {
maxConnToSet = user.MaxConn
cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID)
}
for _, fp := range ports {"""
content = content.replace(resolution, resolution_new)
# Inject ensureConnLimiterOnNode inside loop
loop_inner = """ if limiterID != nil && speed != nil {
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
// If the limiter push fails because the node is offline, skip it with a warning
if isNodeOfflineOrTimeoutError(err) {
node, _ := h.getNodeRecord(fp.NodeID)
nodeName := fmt.Sprintf("%d", fp.NodeID)
if node != nil && strings.TrimSpace(node.Name) != "" {
nodeName = strings.TrimSpace(node.Name)
}
warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName))
continue
}
return nil, err
}
}
node, err := h.getNodeRecord(fp.NodeID)"""
loop_inner_new = """ if limiterID != nil && speed != nil {
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
// If the limiter push fails because the node is offline, skip it with a warning
if isNodeOfflineOrTimeoutError(err) {
node, _ := h.getNodeRecord(fp.NodeID)
nodeName := fmt.Sprintf("%d", fp.NodeID)
if node != nil && strings.TrimSpace(node.Name) != "" {
nodeName = strings.TrimSpace(node.Name)
}
warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName))
continue
}
return nil, err
}
}
if cLimiterName != "" {
if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil {
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
}
}
node, err := h.getNodeRecord(fp.NodeID)"""
content = content.replace(loop_inner, loop_inner_new)
# Update buildForwardServiceConfigs call in syncForwardServicesWithWarnings
content = content.replace(
'services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID)',
'services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName)'
)
# Update fallbackForwardPortToDefaultBind call
content = content.replace(
'warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID)',
'warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID, cLimiterName)'
)
# 5. Update fallbackForwardPortToDefaultBind declaration and logic
content = content.replace(
'func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64) (string, error) {',
'func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64, cLimiterName string) (string, error) {'
)
content = content.replace(
'defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID)',
'defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName)'
)
with open('internal/http/handler/control_plane.go', 'w') as f:
f.write(content)
# Update control_plane_test.go
with open('internal/http/handler/control_plane_test.go', 'r') as f:
test_content = f.read()
test_content = re.sub(
r'buildForwardServiceConfigs\((.*?),(.*?),(.*?),(.*?),(.*?),(.*?),(.*?)\)',
r'buildForwardServiceConfigs(\1,\2,\3,\4,\5,\6,\7, "")',
test_content
)
with open('internal/http/handler/control_plane_test.go', 'w') as f:
f.write(test_content)
@@ -0,0 +1,107 @@
package contract_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"go-backend/internal/store/model"
)
func TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately(t *testing.T) {
secret := "monitoring-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now()
nowMs := now.UnixMilli()
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
monthKey := int64(now.Year()*100 + int(now.Month()))
const bytesPerGB = int64(1024 * 1024 * 1024)
node := &model.Node{Name: "node-1", Secret: "node-secret", ServerIP: "127.0.0.1", Port: "10000-10010", TCPListenAddr: "[::]", UDPListenAddr: "[::]", CreatedTime: nowMs, Status: 1}
if err := repo.DB().Create(node).Error; err != nil {
t.Fatalf("seed node: %v", err)
}
if err := repo.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'flow_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
tunnel := &model.Tunnel{Name: "tunnel-1", TrafficRatio: 1.0, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
if err := repo.DB().Create(tunnel).Error; err != nil {
t.Fatalf("seed tunnel: %v", err)
}
if err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnel.ID).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
forward := &model.Forward{ID: 20, UserID: 2, UserName: "flow_user", Name: "forward-20", TunnelID: tunnel.ID, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
if err := repo.DB().Create(forward).Error; err != nil {
t.Fatalf("seed forward: %v", err)
}
if err := repo.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 1, 0, ?, ?, ?, ?, 0, 0, '', ?, ?)`, bytesPerGB-100, bytesPerGB-100, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user_quota: %v", err)
}
body, err := json.Marshal([]map[string]interface{}{
{"n": "20_2_10", "u": 70, "d": 50},
{"n": "20_2_10_tcp", "u": 40, "d": 30},
})
if err != nil {
t.Fatalf("marshal body: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
if res.Code != http.StatusOK {
t.Fatalf("expected status 200, got %d", res.Code)
}
if got := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = 20`); got != 0 {
t.Fatalf("expected forward paused immediately, got status=%d", got)
}
if got := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`); got != 1 {
t.Fatalf("expected quota disabled flag=1, got %d", got)
}
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 80 {
t.Fatalf("expected forward in_flow=80, got %d", got)
}
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 110 {
t.Fatalf("expected forward out_flow=110, got %d", got)
}
metrics, err := repo.GetTunnelMetrics(tunnel.ID, 0, nowMs+60_000)
if err != nil {
t.Fatalf("get tunnel metrics: %v", err)
}
if len(metrics) != 1 || metrics[0].BytesIn != 80 || metrics[0].BytesOut != 110 {
t.Fatalf("expected one aggregated metric row, got %#v", metrics)
}
body, err = json.Marshal([]map[string]interface{}{
{"n": "20_2_10", "u": 10, "d": 20},
{"n": "20_2_10", "u": 10, "d": 20},
{"n": "20_2_10_tcp", "u": 10, "d": 20},
})
if err != nil {
t.Fatalf("marshal body: %v", err)
}
req = httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
res = httptest.NewRecorder()
router.ServeHTTP(res, req)
if res.Code != http.StatusOK {
t.Fatalf("expected second request status 200, got %d", res.Code)
}
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 140 {
t.Fatalf("expected forward in_flow=140 after second request, got %d", got)
}
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 140 {
t.Fatalf("expected forward out_flow=140 after second request, got %d", got)
}
metrics, err = repo.GetTunnelMetrics(tunnel.ID, 0, nowMs+60_000)
if err != nil {
t.Fatalf("get tunnel metrics after second request: %v", err)
}
if len(metrics) != 1 || metrics[0].BytesIn != 140 || metrics[0].BytesOut != 140 {
t.Fatalf("expected one aggregated metric row after second request, got %#v", metrics)
}
}
@@ -1085,6 +1085,179 @@ func jsonNumber(v int64) string {
return strconv.FormatInt(v, 10)
}
func TestForwardIPSpeedLimitPermission(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now().UnixMilli()
if err := repo.DB().Exec(`
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'normal_user', 'pwd', 1, ?, 99999, 0, 0, 1, 10, ?, ?, 1)
`, now+86400000, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(12, 'ip-speed-permission-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(20, 'ip-speed-permission-node', 'ip-speed-permission-secret', '10.22.0.1', '10.22.0.1', '', '32200-32210', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
`, now, now).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(12, 1, 20, 32201, 'round', 1, 'tls')
`).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO speed_limit(id, name, speed, created_time, status)
VALUES(9, 'per-ip-10m', 10, ?, 1)
`, now).Error; err != nil {
t.Fatalf("insert speed limit: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(2, 12, 10, 99999, 0, 0, 1, ?, 1)
`, now+86400000).Error; err != nil {
t.Fatalf("insert user tunnel: %v", err)
}
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
if err != nil {
t.Fatalf("generate user token: %v", err)
}
body, err := json.Marshal(map[string]interface{}{
"name": "blocked-ip-speed",
"tunnelId": 12,
"remoteAddr": "1.1.1.1:443",
"strategy": "fifo",
"ipSpeedId": 9,
})
if err != nil {
t.Fatalf("marshal create payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(body))
req.Header.Set("Authorization", userToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
assertCodeMsg(t, res, -1, "普通用户无法设置每 IP 限速规则")
}
func TestForwardIPSpeedLimitUpdatePermission(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
server := httptest.NewServer(router)
defer server.Close()
now := time.Now().UnixMilli()
if err := repo.DB().Exec(`
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'normal_user_ip_update', 'pwd', 1, ?, 99999, 0, 0, 1, 10, ?, ?, 1)
`, now+86400000, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(13, 'ip-speed-update-permission-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(21, 'ip-speed-update-permission-node', 'ip-speed-update-permission-secret', '10.22.0.2', '10.22.0.2', '', '32300-32310', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
`, now, now).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(13, 1, 21, 32301, 'round', 1, 'tls')
`).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO speed_limit(id, name, speed, created_time, status)
VALUES(10, 'per-ip-10m-update', 10, ?, 1), (11, 'per-ip-20m-update', 20, ?, 1)
`, now, now).Error; err != nil {
t.Fatalf("insert speed limits: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(2, 13, 10, 99999, 0, 0, 1, ?, 1)
`, now+86400000).Error; err != nil {
t.Fatalf("insert user tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, ip_speed_id, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(30, 2, 'normal_user_ip_update', 'ip-speed-update-forward', 13, '1.1.1.1:443', 'fifo', 10, 0, 0, ?, ?, 1, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
userToken, err := auth.GenerateToken(2, "normal_user_ip_update", 1, secret)
if err != nil {
t.Fatalf("generate user token: %v", err)
}
stopNode := startMockNodeSession(t, server.URL, "ip-speed-update-permission-secret")
defer stopNode()
updateForward := func(t *testing.T, ipSpeedID interface{}) *httptest.ResponseRecorder {
t.Helper()
if err := repo.DB().Exec(`UPDATE forward SET ip_speed_id = 10 WHERE id = 30`).Error; err != nil {
t.Fatalf("reset forward ip speed limit: %v", err)
}
body, err := json.Marshal(map[string]interface{}{
"id": 30,
"name": "ip-speed-update-forward",
"tunnelId": 13,
"remoteAddr": "1.1.1.1:443",
"ipSpeedId": ipSpeedID,
})
if err != nil {
t.Fatalf("marshal update payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(body))
req.Header.Set("Authorization", userToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
return res
}
assertStoredIPSpeedID := func(t *testing.T, want int64) {
t.Helper()
var got sql.NullInt64
if err := repo.DB().Raw(`SELECT ip_speed_id FROM forward WHERE id = 30`).Scan(&got).Error; err != nil {
t.Fatalf("read forward ip_speed_id: %v", err)
}
if !got.Valid || got.Int64 != want {
t.Fatalf("expected ip_speed_id %d, got valid=%v value=%d", want, got.Valid, got.Int64)
}
}
t.Run("non-admin cannot change existing ipSpeedId", func(t *testing.T) {
res := updateForward(t, 11)
assertCodeMsg(t, res, -1, "普通用户无法修改每 IP 限速规则")
assertStoredIPSpeedID(t, 10)
})
t.Run("non-admin cannot clear existing ipSpeedId", func(t *testing.T) {
res := updateForward(t, nil)
assertCodeMsg(t, res, -1, "普通用户无法修改每 IP 限速规则")
assertStoredIPSpeedID(t, 10)
})
t.Run("non-admin can keep existing ipSpeedId", func(t *testing.T) {
res := updateForward(t, 10)
assertCode(t, res, 0)
assertStoredIPSpeedID(t, 10)
})
}
func TestNonAdminCannotSetSpeedIdOrPort(t *testing.T) {
secret := "contract-jwt-secret-perm"
router, repo := setupContractRouter(t, secret)
@@ -0,0 +1,219 @@
package contract_test
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"go-backend/internal/auth"
"go-backend/internal/http/handler"
"go-backend/internal/http/response"
)
func TestForwardLocalRemoteAddrToggleContracts(t *testing.T) {
handler.DisableSafeRemoteAddrCheckForTesting = false
t.Cleanup(func() {
handler.DisableSafeRemoteAddrCheckForTesting = true
})
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
server := httptest.NewServer(router)
defer server.Close()
now := time.Now().UnixMilli()
if err := repo.DB().Exec(`
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'local_remote_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
`, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "local-remote-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
tunnelID := mustLastInsertID(t, repo, "local-remote-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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "local-remote-entry", "local-remote-secret", "10.60.0.1", "10.60.0.1", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
entryNodeID := mustLastInsertID(t, repo, "local-remote-entry")
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, entryNodeID).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(601, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
`, tunnelID).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
userToken, err := auth.GenerateToken(2, "local_remote_user", 1, secret)
if err != nil {
t.Fatalf("generate user token: %v", err)
}
stopNode := startMockNodeSession(t, server.URL, "local-remote-secret")
defer stopNode()
waitNodeStatus(t, repo, entryNodeID, 1)
t.Run("local remote address is rejected on create when toggle is off", func(t *testing.T) {
createPayload := map[string]interface{}{
"name": "deny-local-create",
"tunnelId": tunnelID,
"remoteAddr": "127.0.0.1:8080",
"strategy": "fifo",
}
createBody, err := json.Marshal(createPayload)
if err != nil {
t.Fatalf("marshal create payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
req.Header.Set("Authorization", userToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
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 local remote address to be rejected when toggle is off")
}
if !strings.Contains(out.Msg, "internal IP") && !strings.Contains(out.Msg, "内部") {
t.Fatalf("expected internal IP error, got code=%d msg=%q", out.Code, out.Msg)
}
})
t.Run("local remote address is allowed on create when toggle is on", func(t *testing.T) {
if err := repo.DB().Exec(`
INSERT INTO vite_config(name, value, time)
VALUES(?, ?, ?)
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
`, "allow_local_remote_addr", "true", time.Now().UnixMilli()).Error; err != nil {
t.Fatalf("enable allow_local_remote_addr: %v", err)
}
createPayload := map[string]interface{}{
"name": "allow-local-create",
"tunnelId": tunnelID,
"remoteAddr": "127.0.0.1:8080",
"strategy": "fifo",
}
createBody, err := json.Marshal(createPayload)
if err != nil {
t.Fatalf("marshal create payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
req.Header.Set("Authorization", userToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
assertCode(t, res, 0)
})
t.Run("local remote address is rejected on update when toggle is off", func(t *testing.T) {
if err := repo.DB().Exec(`
INSERT INTO vite_config(name, value, time)
VALUES(?, ?, ?)
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
`, "allow_local_remote_addr", "false", time.Now().UnixMilli()).Error; err != nil {
t.Fatalf("disable allow_local_remote_addr: %v", err)
}
createPayload := map[string]interface{}{
"name": "safe-remote-before-update",
"tunnelId": tunnelID,
"remoteAddr": "8.8.8.8:53",
"strategy": "fifo",
}
createBody, err := json.Marshal(createPayload)
if err != nil {
t.Fatalf("marshal safe create payload: %v", err)
}
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
createReq.Header.Set("Authorization", userToken)
createReq.Header.Set("Content-Type", "application/json")
createRes := httptest.NewRecorder()
router.ServeHTTP(createRes, createReq)
assertCode(t, createRes, 0)
forwardID := mustLastInsertID(t, repo, "safe-remote-before-update")
updatePayload := map[string]interface{}{
"id": forwardID,
"name": "safe-remote-before-update",
"tunnelId": tunnelID,
"remoteAddr": "127.0.0.1:8081",
"strategy": "fifo",
}
updateBody, err := json.Marshal(updatePayload)
if err != nil {
t.Fatalf("marshal update payload: %v", err)
}
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
updateReq.Header.Set("Authorization", userToken)
updateReq.Header.Set("Content-Type", "application/json")
updateRes := httptest.NewRecorder()
router.ServeHTTP(updateRes, updateReq)
var out response.R
if err := json.NewDecoder(updateRes.Body).Decode(&out); err != nil {
t.Fatalf("decode update response: %v", err)
}
if out.Code == 0 {
t.Fatalf("expected local remote address to be rejected on update when toggle is off")
}
if !strings.Contains(out.Msg, "internal IP") && !strings.Contains(out.Msg, "内部") {
t.Fatalf("expected internal IP error on update, got code=%d msg=%q", out.Code, out.Msg)
}
})
t.Run("local remote address is allowed on update when toggle is on", func(t *testing.T) {
if err := repo.DB().Exec(`
INSERT INTO vite_config(name, value, time)
VALUES(?, ?, ?)
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
`, "allow_local_remote_addr", "true", time.Now().UnixMilli()).Error; err != nil {
t.Fatalf("enable allow_local_remote_addr: %v", err)
}
var forwardID int64
if err := repo.DB().Raw(`SELECT id FROM forward WHERE name = ? ORDER BY id DESC LIMIT 1`, "safe-remote-before-update").Row().Scan(&forwardID); err != nil {
t.Fatalf("query forward id: %v", err)
}
updatePayload := map[string]interface{}{
"id": forwardID,
"name": "safe-remote-before-update",
"tunnelId": tunnelID,
"remoteAddr": "127.0.0.1:8081",
"strategy": "fifo",
}
updateBody, err := json.Marshal(updatePayload)
if err != nil {
t.Fatalf("marshal update payload: %v", err)
}
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
updateReq.Header.Set("Authorization", userToken)
updateReq.Header.Set("Content-Type", "application/json")
updateRes := httptest.NewRecorder()
router.ServeHTTP(updateRes, updateReq)
assertCode(t, updateRes, 0)
})
}
@@ -0,0 +1,605 @@
package contract_test
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"reflect"
"strings"
"sync"
"testing"
"time"
"github.com/gorilla/websocket"
"go-backend/internal/auth"
"go-backend/internal/http/response"
"go-backend/internal/security"
)
func TestMaxConnLimit(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "max-conn-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT id FROM tunnel WHERE name = ?", "max-conn-tunnel").Scan(&tunnelID).Error; err != nil {
t.Fatalf("get tunnel ID: %v", err)
}
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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "max-conn-node", "max-conn-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)
}
var nodeID int64
if err := r.DB().Raw("SELECT id FROM node WHERE name = ?", "max-conn-node").Scan(&nodeID).Error; err != nil {
t.Fatalf("get node ID: %v", err)
}
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 user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(1, ?, 10, 99999, 0, 0, 1, ?, 1)
`, tunnelID, now+365*24*3600*1000).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
var commandMu sync.Mutex
receivedCommands := make([]string, 0)
var addCLimitersData json.RawMessage
var updateCLimitersData json.RawMessage
stopNode := startMockSessionForMaxConn(t, server.URL, "max-conn-secret", func(cmdType string, data json.RawMessage) (bool, string) {
commandMu.Lock()
defer commandMu.Unlock()
receivedCommands = append(receivedCommands, cmdType)
if cmdType == "AddCLimiters" {
addCLimitersData = append([]byte(nil), data...)
return true, "already exists"
}
if cmdType == "UpdateCLimiters" {
updateCLimitersData = append([]byte(nil), data...)
}
return false, ""
})
defer stopNode()
waitNodeStatus(t, r, nodeID, 1)
payload := map[string]interface{}{
"name": "max-conn-forward",
"tunnelId": tunnelID,
"remoteAddr": "1.1.1.1:443",
"strategy": "fifo",
"maxConn": 42,
"ipMaxConn": 7,
"proxyProtocol": 2,
}
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 success, got code=%d msg=%s", out.Code, out.Msg)
}
var forwardID int64
if err := r.DB().Raw("SELECT id FROM forward WHERE name = ?", "max-conn-forward").Scan(&forwardID).Error; err != nil {
t.Fatalf("get forward ID: %v", err)
}
listOut := requestContractEnvelope(t, router, adminToken, "/api/v1/forward/list", nil)
if listOut.Code != 0 {
t.Fatalf("expected /forward/list success, got code=%d msg=%s", listOut.Code, listOut.Msg)
}
rows := mustContractSlice(t, listOut.Data, "forward list")
var target map[string]interface{}
for _, row := range rows {
item, ok := row.(map[string]interface{})
if !ok {
t.Fatalf("expected forward item to be object, got %T", row)
}
idVal, ok := item["id"].(float64)
if !ok {
t.Fatalf("expected forward id to be float64, got %T", item["id"])
}
if int64(idVal) == forwardID {
target = item
break
}
}
if target == nil {
t.Fatalf("forward %d not found in /forward/list response", forwardID)
}
maxConnVal, ok := target["maxConn"].(float64)
if !ok {
t.Fatalf("expected maxConn to be float64, got %T (%v)", target["maxConn"], target["maxConn"])
}
if int(maxConnVal) != 42 {
t.Fatalf("expected maxConn 42 in /forward/list, got %v", maxConnVal)
}
proxyProtocolVal, ok := target["proxyProtocol"].(float64)
if !ok {
t.Fatalf("expected proxyProtocol to be float64, got %T (%v)", target["proxyProtocol"], target["proxyProtocol"])
}
if int(proxyProtocolVal) != 2 {
t.Fatalf("expected proxyProtocol 2 in /forward/list, got %v", proxyProtocolVal)
}
commandMu.Lock()
defer commandMu.Unlock()
hasAdd := false
hasUpdate := false
for _, cmd := range receivedCommands {
if cmd == "AddCLimiters" {
hasAdd = true
}
if cmd == "UpdateCLimiters" {
hasUpdate = true
}
}
if !hasAdd {
t.Fatalf("expected AddCLimiters to be sent, but it was not. Received: %v", receivedCommands)
}
if !hasUpdate {
t.Fatalf("expected UpdateCLimiters to be sent after AddCLimiters failed with already exists. Received: %v", receivedCommands)
}
expectedName := fmt.Sprintf("rule_conn_limit_%d", forwardID)
// verify payload for AddCLimiters
var addData map[string]interface{}
if err := json.Unmarshal(addCLimitersData, &addData); err != nil {
t.Fatalf("unmarshal AddCLimiters data: %v", err)
}
if addData["name"] != expectedName {
t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"])
}
if limits, ok := addData["limits"].([]interface{}); ok {
if len(limits) != 2 || limits[0] != "$ 42" || limits[1] != "$$ 7" {
t.Fatalf("expected limits to contain '$ 42' and '$$ 7', got %v", limits)
}
} else {
t.Fatalf("invalid limits type in AddCLimiters data: %v", addData)
}
// verify payload for UpdateCLimiters
var updateData map[string]interface{}
if err := json.Unmarshal(updateCLimitersData, &updateData); err != nil {
t.Fatalf("unmarshal UpdateCLimiters data: %v", err)
}
if updateData["limiter"] != expectedName {
t.Fatalf("expected update limiter name %s, got %v", expectedName, updateData["limiter"])
}
nestedData, ok := updateData["data"].(map[string]interface{})
if !ok {
t.Fatalf("expected nested 'data' in UpdateCLimiters, got %v", updateData)
}
if nestedData["name"] != expectedName {
t.Fatalf("expected nested name %s, got %v", expectedName, nestedData["name"])
}
if nestedLimits, ok := nestedData["limits"].([]interface{}); ok {
if len(nestedLimits) != 2 || nestedLimits[0] != "$ 42" || nestedLimits[1] != "$$ 7" {
t.Fatalf("expected nested limits to contain '$ 42' and '$$ 7', got %v", nestedLimits)
}
} else {
t.Fatalf("invalid limits type in UpdateCLimiters nested data: %v", nestedData)
}
}
func TestUserMaxConnUpdateResyncsExistingForwards(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, max_conn, created_time, updated_time, status)
VALUES(2, 'limited_user', 'pwd', 1, ?, 99999, 0, 0, 1, 10, 0, ?, ?, 1)
`, now+365*24*3600*1000, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(10, 'user-max-conn-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(20, 'user-max-conn-node', 'user-max-conn-secret', '10.21.0.1', '10.21.0.1', '', '32100-32110', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
`, now, now).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(10, 1, 20, 32101, 'round', 1, 'tls')
`).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(30, 2, 10, 10, 99999, 0, 0, 1, ?, 1)
`, now+365*24*3600*1000).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx, max_conn)
VALUES(40, 2, 'limited_user', 'user-max-conn-forward', 10, '1.1.1.1:443', 'fifo', 0, 0, ?, ?, 1, 0, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO forward_port(forward_id, node_id, port, in_ip)
VALUES(40, 20, 32105, '')
`).Error; err != nil {
t.Fatalf("insert forward_port: %v", err)
}
var commandMu sync.Mutex
receivedCommands := make([]string, 0)
var addCLimitersData json.RawMessage
var updateServiceData json.RawMessage
stopNode := startMockSessionForMaxConn(t, server.URL, "user-max-conn-secret", func(cmdType string, data json.RawMessage) (bool, string) {
commandMu.Lock()
defer commandMu.Unlock()
receivedCommands = append(receivedCommands, cmdType)
if cmdType == "AddCLimiters" {
addCLimitersData = append([]byte(nil), data...)
}
if cmdType == "UpdateService" {
updateServiceData = append([]byte(nil), data...)
}
return false, ""
})
defer stopNode()
waitNodeStatus(t, r, 20, 1)
payload := map[string]interface{}{
"id": 2,
"user": "limited_user",
"flow": 99999,
"num": 10,
"expTime": now + 365*24*3600*1000,
"flowResetTime": 1,
"status": 1,
"maxConn": 37,
}
body, err := json.Marshal(payload)
if err != nil {
t.Fatalf("marshal payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/update", bytes.NewReader(body))
req.Header.Set("Authorization", adminToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
var out response.R
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
t.Fatalf("decode response: %v", err)
}
if out.Code != 0 {
t.Fatalf("expected user update success, got code=%d msg=%s", out.Code, out.Msg)
}
commandMu.Lock()
defer commandMu.Unlock()
if addCLimitersData == nil {
t.Fatalf("expected AddCLimiters after user maxConn update. Received: %v", receivedCommands)
}
if updateServiceData == nil {
t.Fatalf("expected UpdateService after user maxConn update. Received: %v", receivedCommands)
}
var addData map[string]interface{}
if err := json.Unmarshal(addCLimitersData, &addData); err != nil {
t.Fatalf("unmarshal AddCLimiters data: %v", err)
}
if addData["name"] != "user_conn_limit_2" {
t.Fatalf("expected limiter name user_conn_limit_2, got %v", addData["name"])
}
limits, ok := addData["limits"].([]interface{})
if !ok || len(limits) != 1 || limits[0] != "$ 37" {
t.Fatalf("expected limits to contain '$ 37', got %v", addData["limits"])
}
var services []map[string]interface{}
if err := json.Unmarshal(updateServiceData, &services); err != nil {
t.Fatalf("unmarshal UpdateService data: %v", err)
}
if len(services) != 2 {
t.Fatalf("expected 2 services, got %d", len(services))
}
for _, service := range services {
if service["climiter"] != "user_conn_limit_2" {
t.Fatalf("expected service climiter user_conn_limit_2, got %v", service["climiter"])
}
}
}
func TestUserMaxConnWithPerIPRuleSplitsRuntimeLimiters(t *testing.T) {
secret := "contract-jwt-secret"
router, r := setupContractRouter(t, secret)
server := httptest.NewServer(router)
defer server.Close()
userToken, err := auth.GenerateToken(3, "per_ip_user", 1, secret)
if err != nil {
t.Fatalf("generate user 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, max_conn, created_time, updated_time, status)
VALUES(3, 'per_ip_user', 'pwd', 1, ?, 99999, 0, 0, 1, 10, 37, ?, ?, 1)
`, now+365*24*3600*1000, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(11, 'user-per-ip-conn-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
VALUES(21, 'user-per-ip-conn-node', 'user-per-ip-conn-secret', '10.23.0.1', '10.23.0.1', '', '32300-32310', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
`, now, now).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(11, 1, 21, 32301, 'round', 1, 'tls')
`).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(31, 3, 11, 10, 99999, 0, 0, 1, ?, 1)
`, now+365*24*3600*1000).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
var commandMu sync.Mutex
receivedCommands := make([]string, 0)
addCLimitersData := make([]json.RawMessage, 0)
var updateServiceData json.RawMessage
stopNode := startMockSessionForMaxConn(t, server.URL, "user-per-ip-conn-secret", func(cmdType string, data json.RawMessage) (bool, string) {
commandMu.Lock()
defer commandMu.Unlock()
receivedCommands = append(receivedCommands, cmdType)
if cmdType == "AddCLimiters" {
addCLimitersData = append(addCLimitersData, append([]byte(nil), data...))
}
if cmdType == "UpdateService" {
updateServiceData = append([]byte(nil), data...)
}
return false, ""
})
defer stopNode()
waitNodeStatus(t, r, 21, 1)
payload := map[string]interface{}{
"name": "user-per-ip-conn-forward",
"tunnelId": int64(11),
"remoteAddr": "1.1.1.1:443",
"strategy": "fifo",
"ipMaxConn": 7,
}
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", userToken)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
router.ServeHTTP(res, req)
var out response.R
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
t.Fatalf("decode response: %v", err)
}
if out.Code != 0 {
t.Fatalf("expected create success, got code=%d msg=%s", out.Code, out.Msg)
}
var forwardID int64
if err := r.DB().Raw("SELECT id FROM forward WHERE name = ?", "user-per-ip-conn-forward").Scan(&forwardID).Error; err != nil {
t.Fatalf("get forward ID: %v", err)
}
expectedRuleName := fmt.Sprintf("rule_conn_limit_%d", forwardID)
commandMu.Lock()
defer commandMu.Unlock()
if len(addCLimitersData) != 2 {
t.Fatalf("expected two AddCLimiters commands. Received: %v", receivedCommands)
}
if updateServiceData == nil {
t.Fatalf("expected UpdateService. Received: %v", receivedCommands)
}
gotLimits := make(map[string][]string)
for _, raw := range addCLimitersData {
var data map[string]interface{}
if err := json.Unmarshal(raw, &data); err != nil {
t.Fatalf("unmarshal AddCLimiters data: %v", err)
}
limits, ok := data["limits"].([]interface{})
if !ok {
t.Fatalf("expected limits array, got %T", data["limits"])
}
for _, limit := range limits {
gotLimits[fmt.Sprint(data["name"])] = append(gotLimits[fmt.Sprint(data["name"])], fmt.Sprint(limit))
}
}
if !reflect.DeepEqual(gotLimits["user_conn_limit_3"], []string{"$ 37"}) {
t.Fatalf("expected user max limiter payload, got %v", gotLimits["user_conn_limit_3"])
}
if !reflect.DeepEqual(gotLimits[expectedRuleName], []string{"$$ 7"}) {
t.Fatalf("expected rule per-IP limiter payload, got %v", gotLimits[expectedRuleName])
}
var services []map[string]interface{}
if err := json.Unmarshal(updateServiceData, &services); err != nil {
t.Fatalf("unmarshal UpdateService data: %v", err)
}
expectedCLimiter := "user_conn_limit_3," + expectedRuleName
for _, service := range services {
if service["climiter"] != expectedCLimiter {
t.Fatalf("expected service climiter %s, got %v", expectedCLimiter, service["climiter"])
}
}
}
func startMockSessionForMaxConn(t *testing.T, baseURL string, nodeSecret string, onCommand func(cmdType string, data json.RawMessage) (bool, 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"`
Data json.RawMessage `json:"data"`
}
if err := json.Unmarshal(plain, &cmd); err != nil {
continue
}
if strings.TrimSpace(cmd.RequestID) == "" {
continue
}
shouldFail := false
failMsg := ""
if onCommand != nil {
shouldFail, failMsg = onCommand(strings.TrimSpace(cmd.Type), cmd.Data)
}
respType := fmt.Sprintf("%sResponse", cmd.Type)
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()
})
}
}
@@ -349,9 +349,9 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
tunnelID := mustLastInsertID(t, r, "backup-forward-tunnel")
if err := r.DB().Exec(`
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88).Error; err != nil {
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx, proxy_protocol)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88, 2).Error; err != nil {
t.Fatalf("seed forward for backup: %v", err)
}
forwardID := mustLastInsertID(t, r, "backup-forward")
@@ -412,6 +412,9 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
if !ok {
t.Fatalf("expected forwardPorts for forward %d in payload", forwardID)
}
if proxyProtocol, ok := forwardMap["proxyProtocol"].(float64); !ok || int(proxyProtocol) != 2 {
t.Fatalf("expected exported proxyProtocol 2 for forward %d, got %v", forwardID, forwardMap["proxyProtocol"])
}
for _, p := range portsRaw {
portMap, ok := p.(map[string]interface{})
if !ok {
@@ -475,6 +478,14 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
t.Fatalf("expected forward_port node=%d port=%d after import, got %v", nodeID, port, after)
}
}
var proxyProtocol int
if err := r.DB().Raw(`SELECT proxy_protocol FROM forward WHERE id = ?`, forwardID).Row().Scan(&proxyProtocol); err != nil {
t.Fatalf("query proxy_protocol after import: %v", err)
}
if proxyProtocol != 2 {
t.Fatalf("expected proxy_protocol 2 after import, got %d", proxyProtocol)
}
})
t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) {
@@ -0,0 +1,176 @@
package contract_test
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"reflect"
"sync"
"testing"
"time"
"go-backend/internal/auth"
"go-backend/internal/http/response"
)
func TestPerIPSpeedLimitRuntimePayload(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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "per-ip-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
var tunnelID int64
if err := r.DB().Raw("SELECT id FROM tunnel WHERE name = ?", "per-ip-speed-tunnel").Scan(&tunnelID).Error; err != nil {
t.Fatalf("get tunnel ID: %v", err)
}
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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, "per-ip-speed-node", "per-ip-speed-secret", "10.22.0.1", "10.22.0.1", "", "32200-32210", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
t.Fatalf("insert node: %v", err)
}
var nodeID int64
if err := r.DB().Raw("SELECT id FROM node WHERE name = ?", "per-ip-speed-node").Scan(&nodeID).Error; err != nil {
t.Fatalf("get node ID: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
VALUES(?, 1, ?, 32201, 'round', 1, 'tls')
`, tunnelID, nodeID).Error; err != nil {
t.Fatalf("insert chain_tunnel: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO user_tunnel(user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(1, ?, 10, 99999, 0, 0, 1, ?, 1)
`, tunnelID, now+365*24*3600*1000).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
totalSpeedID, err := r.CreateSpeedLimit("per-ip-total-speed", 80, now, 1)
if err != nil {
t.Fatalf("create total speed limit: %v", err)
}
ipSpeedID, err := r.CreateSpeedLimit("per-ip-client-speed", 40, now, 1)
if err != nil {
t.Fatalf("create per-ip speed limit: %v", err)
}
var commandMu sync.Mutex
receivedCommands := make([]string, 0)
addLimitersData := make([]json.RawMessage, 0)
var updateServiceData json.RawMessage
stopNode := startMockSessionForMaxConn(t, server.URL, "per-ip-speed-secret", func(cmdType string, data json.RawMessage) (bool, string) {
commandMu.Lock()
defer commandMu.Unlock()
receivedCommands = append(receivedCommands, cmdType)
if cmdType == "AddLimiters" {
addLimitersData = append(addLimitersData, append([]byte(nil), data...))
}
if cmdType == "UpdateService" {
updateServiceData = append([]byte(nil), data...)
}
return false, ""
})
defer stopNode()
waitNodeStatus(t, r, nodeID, 1)
payload := map[string]interface{}{
"name": "per-ip-speed-forward",
"tunnelId": tunnelID,
"remoteAddr": "1.1.1.1:443",
"strategy": "fifo",
"speedId": totalSpeedID,
"ipSpeedId": ipSpeedID,
}
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 success, got code=%d msg=%s", out.Code, out.Msg)
}
var forwardID int64
if err := r.DB().Raw("SELECT id FROM forward WHERE name = ?", "per-ip-speed-forward").Scan(&forwardID).Error; err != nil {
t.Fatalf("get forward ID: %v", err)
}
expectedName := fmt.Sprintf("rule_traffic_limit_%d", forwardID)
expectedTotalName := fmt.Sprint(totalSpeedID)
expectedRuleLimits := []string{"0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
expectedTotalLimits := []string{"$ 10.0MB 10.0MB"}
commandMu.Lock()
defer commandMu.Unlock()
if len(addLimitersData) != 2 {
t.Fatalf("expected AddLimiters to be sent. Received: %v", receivedCommands)
}
if updateServiceData == nil {
t.Fatalf("expected UpdateService to be sent. Received: %v", receivedCommands)
}
gotLimiterLimits := make(map[string][]string)
for _, raw := range addLimitersData {
var addData map[string]interface{}
if err := json.Unmarshal(raw, &addData); err != nil {
t.Fatalf("unmarshal AddLimiters data: %v", err)
}
name := fmt.Sprint(addData["name"])
limits, ok := addData["limits"].([]interface{})
if !ok {
t.Fatalf("expected limits array, got %T", addData["limits"])
}
gotLimits := make([]string, 0, len(limits))
for _, limit := range limits {
gotLimits = append(gotLimits, fmt.Sprint(limit))
}
gotLimiterLimits[name] = gotLimits
}
if !reflect.DeepEqual(gotLimiterLimits[expectedTotalName], expectedTotalLimits) {
t.Fatalf("expected total limits %v, got %v", expectedTotalLimits, gotLimiterLimits[expectedTotalName])
}
if !reflect.DeepEqual(gotLimiterLimits[expectedName], expectedRuleLimits) {
t.Fatalf("expected rule limits %v, got %v", expectedRuleLimits, gotLimiterLimits[expectedName])
}
var services []map[string]interface{}
if err := json.Unmarshal(updateServiceData, &services); err != nil {
t.Fatalf("unmarshal UpdateService data: %v", err)
}
if len(services) == 0 {
t.Fatalf("expected services in UpdateService")
}
for _, service := range services {
expectedLimiter := expectedTotalName + "," + expectedName
if service["limiter"] != expectedLimiter {
t.Fatalf("expected service limiter %s, got %v", expectedLimiter, service["limiter"])
}
}
}
@@ -0,0 +1,59 @@
package contract_test
import (
"testing"
"time"
"go-backend/internal/auth"
)
func TestUserListReturnsMaxConn(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now().UnixMilli()
if err := repo.DB().Exec(`
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, max_conn, created_time, updated_time, status)
VALUES(2, 'max_conn_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 10, 37, ?, ?, 1)
`, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
if err != nil {
t.Fatalf("generate admin token: %v", err)
}
out := requestContractEnvelope(t, router, adminToken, "/api/v1/user/list", map[string]interface{}{})
if out.Code != 0 {
t.Fatalf("expected /user/list success, got code=%d msg=%s", out.Code, out.Msg)
}
rows := mustContractSlice(t, out.Data, "user list")
var target map[string]interface{}
for _, row := range rows {
item, ok := row.(map[string]interface{})
if !ok {
t.Fatalf("expected user item to be object, got %T", row)
}
idVal, ok := item["id"].(float64)
if !ok {
t.Fatalf("expected user id to be float64, got %T", item["id"])
}
if int64(idVal) == 2 {
target = item
break
}
}
if target == nil {
t.Fatalf("user 2 not found in /user/list response")
}
maxConnVal, ok := target["maxConn"].(float64)
if !ok {
t.Fatalf("expected maxConn to be float64, got %T (%v)", target["maxConn"], target["maxConn"])
}
if int(maxConnVal) != 37 {
t.Fatalf("expected maxConn 37 in /user/list, got %v", maxConnVal)
}
}
+5
View File
@@ -0,0 +1,5 @@
gost
gost_*
gost-*
*.sha256
*.exe
+4
View File
@@ -116,6 +116,10 @@ func main() {
fmt.Printf("✅ 配置加载成功 - addr: %s\n", config.Addr)
// 设置运行时配置持久化路径
socket.SetConfigPersistPath("gost.json")
// 启用持久化将在 program.Start() 后开启,避免启动加载阶段触发冗余写入
log := xlogger.NewLogger()
logger.SetDefault(log)
+5
View File
@@ -16,6 +16,7 @@ import (
metrics "github.com/go-gost/x/metrics/service"
"github.com/go-gost/x/registry"
xservice "github.com/go-gost/x/service"
"github.com/go-gost/x/socket"
"github.com/judwhite/go-svc"
"net/http"
"os"
@@ -66,6 +67,10 @@ func (p *program) Start() error {
return err
}
// Enable config persistence after initial load so runtime mutations
// (AddService, UpdateService, DeleteService, etc.) are saved to disk.
socket.EnableConfigPersist()
if err := p.run(cfg); err != nil {
return err
}
+9 -2
View File
@@ -44,9 +44,14 @@ func Set(c *Config) {
func OnUpdate(f func(c *Config) error) error {
globalMux.Lock()
defer globalMux.Unlock()
err := f(global)
globalMux.Unlock()
return f(global)
if err == nil {
persist()
}
return err
}
type LogConfig struct {
@@ -573,6 +578,7 @@ func (c *Config) Load() error {
if err := v.ReadInConfig(); err != nil {
return err
}
SetPersistPath(v.ConfigFileUsed())
return v.Unmarshal(c)
}
@@ -590,6 +596,7 @@ func (c *Config) ReadFile(file string) error {
if err := v.ReadInConfig(); err != nil {
return err
}
SetPersistPath(v.ConfigFileUsed())
return v.Unmarshal(c)
}
@@ -0,0 +1,199 @@
package service
import (
"context"
"fmt"
"sort"
"strconv"
"strings"
corelimiter "github.com/go-gost/core/limiter"
connlimiter "github.com/go-gost/core/limiter/conn"
trafficlimiter "github.com/go-gost/core/limiter/traffic"
xtraffic "github.com/go-gost/x/limiter/traffic"
"github.com/go-gost/x/registry"
)
func resolveTrafficLimiter(names string) trafficlimiter.TrafficLimiter {
parts := splitLimiterNames(names)
if len(parts) == 0 {
return nil
}
if len(parts) == 1 {
return resolveSingleTrafficLimiter(parts[0])
}
limiters := make([]trafficlimiter.TrafficLimiter, 0, len(parts))
for _, part := range parts {
if lim := resolveSingleTrafficLimiter(part); lim != nil {
limiters = append(limiters, lim)
}
}
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
return &compositeTrafficLimiter{limiters: limiters}
}
func resolveSingleTrafficLimiter(name string) trafficlimiter.TrafficLimiter {
lim := registry.TrafficLimiterRegistry().Get(name)
if lim != nil {
return lim
}
if val, err := strconv.Atoi(name); err == nil && val > 0 {
return xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %dB %dB", xtraffic.ServiceLimitKey, val, val)),
)
}
return xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %s %s", xtraffic.ServiceLimitKey, name, name)),
)
}
func resolveConnLimiter(names string) connlimiter.ConnLimiter {
parts := splitLimiterNames(names)
if len(parts) == 0 {
return nil
}
if len(parts) == 1 {
return registry.ConnLimiterRegistry().Get(parts[0])
}
limiters := make([]connlimiter.ConnLimiter, 0, len(parts))
for _, part := range parts {
if lim := registry.ConnLimiterRegistry().Get(part); lim != nil {
limiters = append(limiters, lim)
}
}
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
return &compositeConnLimiter{limiters: limiters}
}
func splitLimiterNames(names string) []string {
parts := strings.Split(names, ",")
out := make([]string, 0, len(parts))
for _, part := range parts {
if part = strings.TrimSpace(part); part != "" {
out = append(out, part)
}
}
return out
}
type compositeTrafficLimiter struct {
limiters []trafficlimiter.TrafficLimiter
}
func (l *compositeTrafficLimiter) In(ctx context.Context, key string, opts ...corelimiter.Option) trafficlimiter.Limiter {
limiters := make([]trafficlimiter.Limiter, 0, len(l.limiters))
for _, child := range l.limiters {
if lim := child.In(ctx, key, opts...); lim != nil {
limiters = append(limiters, lim)
}
}
return newCompositeTrafficChildLimiter(limiters)
}
func (l *compositeTrafficLimiter) Out(ctx context.Context, key string, opts ...corelimiter.Option) trafficlimiter.Limiter {
limiters := make([]trafficlimiter.Limiter, 0, len(l.limiters))
for _, child := range l.limiters {
if lim := child.Out(ctx, key, opts...); lim != nil {
limiters = append(limiters, lim)
}
}
return newCompositeTrafficChildLimiter(limiters)
}
type compositeTrafficChildLimiter struct {
limiters []trafficlimiter.Limiter
}
func newCompositeTrafficChildLimiter(limiters []trafficlimiter.Limiter) trafficlimiter.Limiter {
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
sort.Slice(limiters, func(i, j int) bool {
return limiters[i].Limit() < limiters[j].Limit()
})
return &compositeTrafficChildLimiter{limiters: limiters}
}
func (l *compositeTrafficChildLimiter) Wait(ctx context.Context, n int) int {
for _, lim := range l.limiters {
if v := lim.Wait(ctx, n); v < n {
n = v
}
}
return n
}
func (l *compositeTrafficChildLimiter) Limit() int {
if len(l.limiters) == 0 {
return 0
}
return l.limiters[0].Limit()
}
func (l *compositeTrafficChildLimiter) Set(n int) {}
type compositeConnLimiter struct {
limiters []connlimiter.ConnLimiter
}
func (l *compositeConnLimiter) Limiter(key string) connlimiter.Limiter {
limiters := make([]connlimiter.Limiter, 0, len(l.limiters))
for _, child := range l.limiters {
if lim := child.Limiter(key); lim != nil {
limiters = append(limiters, lim)
}
}
return newCompositeConnChildLimiter(limiters)
}
type compositeConnChildLimiter struct {
limiters []connlimiter.Limiter
}
func newCompositeConnChildLimiter(limiters []connlimiter.Limiter) connlimiter.Limiter {
if len(limiters) == 0 {
return nil
}
if len(limiters) == 1 {
return limiters[0]
}
sort.Slice(limiters, func(i, j int) bool {
return limiters[i].Limit() < limiters[j].Limit()
})
return &compositeConnChildLimiter{limiters: limiters}
}
func (l *compositeConnChildLimiter) Allow(n int) (allowed bool) {
var i int
for i = range l.limiters {
if allowed = l.limiters[i].Allow(n); !allowed {
break
}
}
if !allowed && i > 0 && n > 0 {
for _, lim := range l.limiters[:i] {
lim.Allow(-n)
}
}
return allowed
}
func (l *compositeConnChildLimiter) Limit() int {
if len(l.limiters) == 0 {
return 0
}
return l.limiters[0].Limit()
}
@@ -0,0 +1,79 @@
package service
import (
"context"
"io"
"testing"
corelimiter "github.com/go-gost/core/limiter"
corelogger "github.com/go-gost/core/logger"
xconn "github.com/go-gost/x/limiter/conn"
xtraffic "github.com/go-gost/x/limiter/traffic"
xlogger "github.com/go-gost/x/logger"
"github.com/go-gost/x/registry"
)
func TestResolveTrafficLimiterComposesCommaSeparatedNames(t *testing.T) {
const totalName = "test_total_speed_composite"
const ruleName = "test_rule_speed_composite"
registry.TrafficLimiterRegistry().Unregister(totalName)
registry.TrafficLimiterRegistry().Unregister(ruleName)
defer registry.TrafficLimiterRegistry().Unregister(totalName)
defer registry.TrafficLimiterRegistry().Unregister(ruleName)
logger := xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))
if err := registry.TrafficLimiterRegistry().Register(totalName, xtraffic.NewTrafficLimiter(xtraffic.LimitsOption("$ 10B 10B"), xtraffic.LoggerOption(logger))); err != nil {
t.Fatalf("register total limiter: %v", err)
}
if err := registry.TrafficLimiterRegistry().Register(ruleName, xtraffic.NewTrafficLimiter(xtraffic.LimitsOption("0.0.0.0/0 3B 3B"), xtraffic.LoggerOption(logger))); err != nil {
t.Fatalf("register rule limiter: %v", err)
}
lim := resolveTrafficLimiter(totalName + "," + ruleName)
if lim == nil {
t.Fatalf("expected composite traffic limiter")
}
serviceLimiter := lim.In(context.Background(), "192.0.2.1:1000", corelimiter.ScopeOption(corelimiter.ScopeService))
if serviceLimiter == nil || serviceLimiter.Limit() != 10 {
t.Fatalf("expected service-scope total limiter 10, got %#v", serviceLimiter)
}
connLimiter := lim.In(context.Background(), "192.0.2.1:1000", corelimiter.ScopeOption(corelimiter.ScopeConn))
if connLimiter == nil || connLimiter.Limit() != 3 {
t.Fatalf("expected conn-scope per-IP limiter 3, got %#v", connLimiter)
}
}
func TestResolveConnLimiterComposesCommaSeparatedNames(t *testing.T) {
const totalName = "test_total_conn_composite"
const ruleName = "test_rule_conn_composite"
registry.ConnLimiterRegistry().Unregister(totalName)
registry.ConnLimiterRegistry().Unregister(ruleName)
defer registry.ConnLimiterRegistry().Unregister(totalName)
defer registry.ConnLimiterRegistry().Unregister(ruleName)
logger := xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))
if err := registry.ConnLimiterRegistry().Register(totalName, xconn.NewConnLimiter(xconn.LimitsOption("$ 2"), xconn.LoggerOption(logger))); err != nil {
t.Fatalf("register total conn limiter: %v", err)
}
if err := registry.ConnLimiterRegistry().Register(ruleName, xconn.NewConnLimiter(xconn.LimitsOption("$$ 1"), xconn.LoggerOption(logger))); err != nil {
t.Fatalf("register rule conn limiter: %v", err)
}
lim := resolveConnLimiter(totalName + "," + ruleName)
if lim == nil {
t.Fatalf("expected composite conn limiter")
}
clientLimiter := lim.Limiter("192.0.2.1")
if clientLimiter == nil || clientLimiter.Limit() != 1 {
t.Fatalf("expected composite client limiter with strictest limit 1, got %#v", clientLimiter)
}
if !clientLimiter.Allow(1) {
t.Fatalf("expected first connection to be allowed")
}
if clientLimiter.Allow(1) {
t.Fatalf("expected per-IP rule limiter to reject second connection")
}
if !lim.Limiter("192.0.2.2").Allow(1) {
t.Fatalf("expected another client to share total limiter but have independent per-IP capacity")
}
}
+4 -19
View File
@@ -3,7 +3,6 @@ package service
import (
"fmt"
"runtime"
"strconv"
"strings"
"time"
@@ -31,7 +30,6 @@ import (
logger_parser "github.com/go-gost/x/config/parsing/logger"
selector_parser "github.com/go-gost/x/config/parsing/selector"
tls_util "github.com/go-gost/x/internal/util/tls"
xtraffic "github.com/go-gost/x/limiter/traffic"
cache_limiter "github.com/go-gost/x/limiter/traffic/cache"
"github.com/go-gost/x/metadata"
mdutil "github.com/go-gost/x/metadata/util"
@@ -135,7 +133,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
postDown = mdutil.GetStrings(md, parsing.MDKeyPostDown)
ignoreChain = mdutil.GetBool(md, parsing.MDKeyIgnoreChain)
if md.IsExists(parsing.MDKeyEnableStats) {
if md != nil && md.IsExists(parsing.MDKeyEnableStats) {
enableStats = mdutil.GetBool(md, parsing.MDKeyEnableStats)
}
@@ -157,7 +155,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
resetTraffic := true
if cfg.Metadata != nil {
md := metadata.NewMetadata(cfg.Metadata)
if md.IsExists(parsing.MDKeyObserverResetTraffic) {
if md != nil && md.IsExists(parsing.MDKeyObserverResetTraffic) {
resetTraffic = mdutil.GetBool(md, parsing.MDKeyObserverResetTraffic)
}
}
@@ -185,20 +183,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
var trafficLimiter listener.Option
if cfg.Limiter != "" {
lim := registry.TrafficLimiterRegistry().Get(cfg.Limiter)
if lim == nil {
// Try to parse as simple number (bandwidth in bytes/sec)
if val, err := strconv.Atoi(cfg.Limiter); err == nil && val > 0 {
lim = xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %dB %dB", xtraffic.ServiceLimitKey, val, val)),
)
}
if lim == nil {
lim = xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption(fmt.Sprintf("%s %s %s", xtraffic.ServiceLimitKey, cfg.Limiter, cfg.Limiter)),
)
}
}
lim := resolveTrafficLimiter(cfg.Limiter)
trafficLimiter = listener.TrafficLimiterOption(
cache_limiter.NewCachedTrafficLimiter(
lim,
@@ -216,7 +201,7 @@ func ParseService(cfg *config.ServiceConfig) (service.Service, error) {
listener.AuthOption(auth_parser.Info(cfg.Listener.Auth)),
listener.TLSConfigOption(tlsConfig),
listener.AdmissionOption(xadmission.AdmissionGroup(admissions...)),
listener.ConnLimiterOption(registry.ConnLimiterRegistry().Get(cfg.CLimiter)),
listener.ConnLimiterOption(resolveConnLimiter(cfg.CLimiter)),
listener.ServiceOption(cfg.Name),
listener.ProxyProtocolOption(ppv),
listener.StatsOption(pStats),
+94
View File
@@ -0,0 +1,94 @@
package config
import (
"bytes"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync"
)
var (
persistPath string
persistMu sync.Mutex
persistEnable bool
)
// SetPersistPath sets the file path where runtime config changes will be
// automatically persisted. Call this once during agent startup before any
// OnUpdate mutations occur.
func SetPersistPath(path string) {
persistMu.Lock()
defer persistMu.Unlock()
persistPath = path
}
func PersistPath() string {
persistMu.Lock()
defer persistMu.Unlock()
return persistPath
}
// EnablePersist turns on automatic persistence. Call this after the initial
// config has been loaded (e.g. after program.Start) so that startup loading
// does not trigger redundant disk writes.
func EnablePersist() {
persistMu.Lock()
defer persistMu.Unlock()
persistEnable = true
}
// persist writes the current global config to the configured file atomically.
func persist() {
persistMu.Lock()
path := persistPath
enabled := persistEnable
persistMu.Unlock()
if !enabled || path == "" {
return
}
cfg := Global()
if cfg == nil {
return
}
var buf bytes.Buffer
enc := json.NewEncoder(&buf)
enc.SetIndent("", " ")
if err := enc.Encode(cfg); err != nil {
fmt.Printf("⚠️ config persist: marshal failed: %v\n", err)
return
}
// Atomic write: write to temp file then rename
dir := filepath.Dir(path)
tmp, err := os.CreateTemp(dir, ".gost-*.tmp")
if err != nil {
fmt.Printf("⚠️ config persist: create temp file failed: %v\n", err)
return
}
tmpName := tmp.Name()
if _, err := tmp.Write(buf.Bytes()); err != nil {
tmp.Close()
os.Remove(tmpName)
fmt.Printf("⚠️ config persist: write failed: %v\n", err)
return
}
if err := tmp.Close(); err != nil {
os.Remove(tmpName)
fmt.Printf("⚠️ config persist: close temp file failed: %v\n", err)
return
}
if err := os.Rename(tmpName, path); err != nil {
os.Remove(tmpName)
fmt.Printf("⚠️ config persist: rename failed: %v\n", err)
return
}
fmt.Printf("💾 节点配置已持久化到 %s\n", path)
}
+32
View File
@@ -0,0 +1,32 @@
package config
import (
"os"
"path/filepath"
"testing"
)
func TestReadFileSetsPersistPath(t *testing.T) {
originalPath := persistPath
originalEnabled := persistEnable
persistPath = ""
persistEnable = false
t.Cleanup(func() {
persistPath = originalPath
persistEnable = originalEnabled
})
dir := t.TempDir()
configFile := filepath.Join(dir, "custom-gost.yaml")
if err := os.WriteFile(configFile, []byte("services: []\n"), 0o644); err != nil {
t.Fatalf("write config file: %v", err)
}
var cfg Config
if err := cfg.ReadFile(configFile); err != nil {
t.Fatalf("ReadFile: %v", err)
}
if persistPath != configFile {
t.Fatalf("expected persistPath %q, got %q", configFile, persistPath)
}
}
+4
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
kcp_util "github.com/go-gost/x/internal/util/kcp"
mdutil "github.com/go-gost/x/metadata/util"
"github.com/go-gost/x/registry"
"github.com/xtaci/kcp-go/v5"
"github.com/xtaci/smux"
@@ -48,6 +49,9 @@ func (d *kcpDialer) Init(md md.Metadata) (err error) {
}
d.md.config.Init()
if md != nil && md.IsExists("kcp.nc") {
d.md.config.NoCongestion = mdutil.GetInt(md, "kcp.nc")
}
return nil
}
+9
View File
@@ -73,6 +73,9 @@ func (d *kcpDialer) parseMetadata(md mdata.Metadata) (err error) {
if md.IsExists("kcp.sndwnd") {
d.md.config.SndWnd = mdutil.GetInt(md, "kcp.sndwnd")
}
if md.IsExists("kcp.sockbuf") {
d.md.config.SockBuf = mdutil.GetInt(md, "kcp.sockbuf")
}
if md.IsExists("kcp.smuxver") {
d.md.config.SmuxVer = mdutil.GetInt(md, "kcp.smuxver")
}
@@ -85,6 +88,12 @@ func (d *kcpDialer) parseMetadata(md mdata.Metadata) (err error) {
if md.IsExists("kcp.nocomp") {
d.md.config.NoComp = mdutil.GetBool(md, "kcp.nocomp")
}
if md.IsExists("kcp.datashard") {
d.md.config.DataShard = mdutil.GetInt(md, "kcp.datashard")
}
if md.IsExists("kcp.parityshard") {
d.md.config.ParityShard = mdutil.GetInt(md, "kcp.parityshard")
}
}
d.md.handshakeTimeout = mdutil.GetDuration(md, handshakeTimeout)
@@ -16,6 +16,7 @@ import (
"github.com/go-gost/core/recorder"
ctxvalue "github.com/go-gost/x/ctx"
xnet "github.com/go-gost/x/internal/net"
"github.com/go-gost/x/internal/net/proxyproto"
"github.com/go-gost/x/internal/util/forwarder"
"github.com/go-gost/x/internal/util/sniffing"
tls_util "github.com/go-gost/x/internal/util/tls"
@@ -252,6 +253,8 @@ func (h *forwardHandler) Handle(ctx context.Context, conn net.Conn, opts ...hand
}
defer cc.Close()
cc = proxyproto.WrapClientConn(h.md.proxyProtocol, conn.RemoteAddr(), conn.LocalAddr(), cc)
if err := xnet.Transport(conn, cc); err != nil {
if marker := target.Marker(); marker != nil {
marker.Mark()
@@ -14,6 +14,7 @@ import (
type metadata struct {
readTimeout time.Duration
proxyProtocol int
httpKeepalive bool
sniffing bool
@@ -38,6 +39,7 @@ func (h *forwardHandler) parseMetadata(md mdata.Metadata) (err error) {
if h.md.readTimeout <= 0 {
h.md.readTimeout = 15 * time.Second
}
h.md.proxyProtocol = mdutil.GetInt(md, "proxyProtocol")
h.md.httpKeepalive = mdutil.GetBool(md, "http.keepalive")
@@ -0,0 +1,111 @@
package local
import (
"bufio"
"context"
"net"
"testing"
"time"
"github.com/go-gost/core/chain"
"github.com/go-gost/core/handler"
"github.com/go-gost/core/hop"
xlogger "github.com/go-gost/x/logger"
xmd "github.com/go-gost/x/metadata"
proxyproto "github.com/pires/go-proxyproto"
)
type proxyProtocolTestHop struct {
node *chain.Node
}
func (h proxyProtocolTestHop) Select(context.Context, ...hop.SelectOption) *chain.Node {
return h.node
}
func (h proxyProtocolTestHop) Nodes() []*chain.Node {
return []*chain.Node{h.node}
}
type proxyProtocolTestRouter struct{}
func (r proxyProtocolTestRouter) Options() *chain.RouterOptions {
return &chain.RouterOptions{}
}
func (r proxyProtocolTestRouter) Dial(ctx context.Context, network, address string) (net.Conn, error) {
var d net.Dialer
return d.DialContext(ctx, network, address)
}
func (r proxyProtocolTestRouter) Bind(context.Context, string, string, ...chain.BindOption) (net.Listener, error) {
return nil, net.ErrClosed
}
func TestLocalForwardHandlerSendsProxyProtocolToTarget(t *testing.T) {
targetListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen target: %v", err)
}
defer targetListener.Close()
entryListener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen entry: %v", err)
}
defer entryListener.Close()
h := NewHandler(
handler.RouterOption(proxyProtocolTestRouter{}),
handler.LoggerOption(xlogger.Nop()),
)
forwarder := h.(handler.Forwarder)
forwarder.Forward(proxyProtocolTestHop{node: chain.NewNode("target", targetListener.Addr().String())})
if err := h.Init(xmd.NewMetadata(map[string]any{"proxyProtocol": 2})); err != nil {
t.Fatalf("init handler: %v", err)
}
handleErr := make(chan error, 1)
acceptErr := make(chan error, 1)
go func() {
serverConn, err := entryListener.Accept()
if err != nil {
acceptErr <- err
return
}
handleErr <- h.Handle(context.Background(), serverConn)
}()
clientConn, err := net.Dial("tcp", entryListener.Addr().String())
if err != nil {
t.Fatalf("dial entry: %v", err)
}
defer clientConn.Close()
targetConn, err := targetListener.Accept()
if err != nil {
t.Fatalf("accept target: %v", err)
}
defer targetConn.Close()
if err := targetConn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("set target deadline: %v", err)
}
header, err := proxyproto.Read(bufio.NewReader(targetConn))
if err != nil {
t.Fatalf("read proxy protocol header: %v", err)
}
if header.Version != 2 {
t.Fatalf("expected proxy protocol v2, got v%d", header.Version)
}
_ = clientConn.Close()
_ = targetConn.Close()
select {
case err := <-acceptErr:
t.Fatalf("accept entry: %v", err)
case <-handleErr:
case <-time.After(2 * time.Second):
t.Fatal("handler did not return after closing connections")
}
}
+30
View File
@@ -0,0 +1,30 @@
package conn
import (
"io"
"testing"
corelogger "github.com/go-gost/core/logger"
xlogger "github.com/go-gost/x/logger"
)
func TestIPLimitKeyCreatesIndependentLimiters(t *testing.T) {
limiter := NewConnLimiter(
LimitsOption("$$ 1"),
LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)
first := limiter.Limiter("192.0.2.1")
second := limiter.Limiter("192.0.2.2")
if first == nil || second == nil {
t.Fatalf("expected non-nil per-IP limiters")
}
if !first.Allow(1) {
t.Fatalf("expected first IP first connection to be allowed")
}
if first.Allow(1) {
t.Fatalf("expected first IP second connection to be rejected")
}
if !second.Allow(1) {
t.Fatalf("expected second IP first connection to be allowed independently")
}
}
+28
View File
@@ -0,0 +1,28 @@
package traffic
import (
"context"
"io"
"testing"
corelogger "github.com/go-gost/core/logger"
xlogger "github.com/go-gost/x/logger"
)
func TestCIDRLimitCreatesIndependentClientLimiters(t *testing.T) {
limiter := NewTrafficLimiter(
LimitsOption("0.0.0.0/0 2B 2B"),
LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)
first := limiter.In(context.Background(), "192.0.2.1:1000")
second := limiter.In(context.Background(), "192.0.2.2:1000")
if first == nil || second == nil {
t.Fatalf("expected non-nil CIDR client limiters")
}
if first == second {
t.Fatalf("expected different clients to receive independent limiter instances")
}
if first.Limit() != 2 || second.Limit() != 2 {
t.Fatalf("expected both limits to be 2, got %d and %d", first.Limit(), second.Limit())
}
}
+4
View File
@@ -15,6 +15,7 @@ import (
limiter_wrapper "github.com/go-gost/x/limiter/traffic/wrapper"
metrics "github.com/go-gost/x/metrics/wrapper"
stats "github.com/go-gost/x/observer/stats/wrapper"
mdutil "github.com/go-gost/x/metadata/util"
"github.com/go-gost/x/registry"
"github.com/xtaci/kcp-go/v5"
"github.com/xtaci/smux"
@@ -53,6 +54,9 @@ func (l *kcpListener) Init(md md.Metadata) (err error) {
config := l.md.config
config.Init()
if md != nil && md.IsExists("kcp.nc") {
config.NoCongestion = mdutil.GetInt(md, "kcp.nc")
}
var conn net.PacketConn
if config.TCP {
+9
View File
@@ -76,6 +76,9 @@ func (l *kcpListener) parseMetadata(md mdata.Metadata) (err error) {
if md.IsExists("kcp.sndwnd") {
l.md.config.SndWnd = mdutil.GetInt(md, "kcp.sndwnd")
}
if md.IsExists("kcp.sockbuf") {
l.md.config.SockBuf = mdutil.GetInt(md, "kcp.sockbuf")
}
if md.IsExists("kcp.smuxver") {
l.md.config.SmuxVer = mdutil.GetInt(md, "kcp.smuxver")
}
@@ -88,6 +91,12 @@ func (l *kcpListener) parseMetadata(md mdata.Metadata) (err error) {
if md.IsExists("kcp.nocomp") {
l.md.config.NoComp = mdutil.GetBool(md, "kcp.nocomp")
}
if md.IsExists("kcp.datashard") {
l.md.config.DataShard = mdutil.GetInt(md, "kcp.datashard")
}
if md.IsExists("kcp.parityshard") {
l.md.config.ParityShard = mdutil.GetInt(md, "kcp.parityshard")
}
}
l.md.backlog = mdutil.GetInt(md, backlog)
+122 -2
View File
@@ -2,8 +2,11 @@ package udp
import (
"net"
"sync"
"time"
"github.com/go-gost/core/limiter"
conn_limiter "github.com/go-gost/core/limiter/conn"
"github.com/go-gost/core/listener"
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
@@ -70,7 +73,7 @@ func (l *udpListener) Init(md md.Metadata) (err error) {
limiter.NetworkOption(conn.LocalAddr().Network()),
)
l.ln = udp.NewListener(conn, &udp.ListenConfig{
ln := udp.NewListener(conn, &udp.ListenConfig{
Backlog: l.md.backlog,
ReadQueueSize: l.md.readQueueSize,
ReadBufferSize: l.md.readBufferSize,
@@ -78,11 +81,128 @@ func (l *udpListener) Init(md md.Metadata) (err error) {
TTL: l.md.ttl,
Logger: l.logger,
})
l.ln = ln
return
}
func (l *udpListener) Accept() (conn net.Conn, err error) {
return l.ln.Accept()
conn, err = l.ln.Accept()
if err != nil {
return
}
if l.options.ConnLimiter != nil {
host, _, _ := net.SplitHostPort(conn.RemoteAddr().String())
if lim := l.options.ConnLimiter.Limiter(host); lim != nil {
if !lim.Allow(1) {
_ = conn.Close()
return newClosedConn(conn), nil
}
conn = wrapConnLimiter(lim, conn)
}
}
if pc, ok := conn.(net.PacketConn); ok {
conn = limiter_wrapper.WrapUDPConn(
pc,
l.options.TrafficLimiter,
conn.RemoteAddr().String(),
limiter.ScopeOption(limiter.ScopeConn),
limiter.ServiceOption(l.options.Service),
limiter.NetworkOption(conn.LocalAddr().Network()),
limiter.SrcOption(conn.RemoteAddr().String()),
)
}
return
}
type connLimiterConn struct {
net.Conn
net.PacketConn
limiter conn_limiter.Limiter
once sync.Once
}
func wrapConnLimiter(limiter conn_limiter.Limiter, conn net.Conn) net.Conn {
pc, ok := conn.(net.PacketConn)
if !ok {
return conn
}
return &connLimiterConn{
Conn: conn,
PacketConn: pc,
limiter: limiter,
}
}
func (c *connLimiterConn) Close() (err error) {
c.once.Do(func() {
c.limiter.Allow(-1)
err = c.Conn.Close()
})
return
}
func (c *connLimiterConn) LocalAddr() net.Addr {
return c.Conn.LocalAddr()
}
func (c *connLimiterConn) SetDeadline(t time.Time) error {
return c.Conn.SetDeadline(t)
}
func (c *connLimiterConn) SetReadDeadline(t time.Time) error {
return c.Conn.SetReadDeadline(t)
}
func (c *connLimiterConn) SetWriteDeadline(t time.Time) error {
return c.Conn.SetWriteDeadline(t)
}
type closedConn struct {
net.Conn
net.PacketConn
}
func newClosedConn(conn net.Conn) net.Conn {
pc, _ := conn.(net.PacketConn)
return closedConn{Conn: conn, PacketConn: pc}
}
func (c closedConn) Read([]byte) (int, error) {
return 0, net.ErrClosed
}
func (c closedConn) Write([]byte) (int, error) {
return 0, net.ErrClosed
}
func (c closedConn) ReadFrom([]byte) (int, net.Addr, error) {
return 0, nil, net.ErrClosed
}
func (c closedConn) WriteTo([]byte, net.Addr) (int, error) {
return 0, net.ErrClosed
}
func (c closedConn) Close() error {
return c.Conn.Close()
}
func (c closedConn) LocalAddr() net.Addr {
return c.Conn.LocalAddr()
}
func (c closedConn) SetDeadline(t time.Time) error {
return c.Conn.SetDeadline(t)
}
func (c closedConn) SetReadDeadline(t time.Time) error {
return c.Conn.SetReadDeadline(t)
}
func (c closedConn) SetWriteDeadline(t time.Time) error {
return c.Conn.SetWriteDeadline(t)
}
func (l *udpListener) Addr() net.Addr {
+160
View File
@@ -0,0 +1,160 @@
package udp
import (
"io"
"net"
"testing"
"time"
corelistener "github.com/go-gost/core/listener"
corelogger "github.com/go-gost/core/logger"
xconn "github.com/go-gost/x/limiter/conn"
xtraffic "github.com/go-gost/x/limiter/traffic"
xlogger "github.com/go-gost/x/logger"
)
func TestAcceptWithLimitersPreservesPacketConn(t *testing.T) {
ln := NewListener(
corelistener.AddrOption("127.0.0.1:0"),
corelistener.ConnLimiterOption(xconn.NewConnLimiter(
xconn.LimitsOption("$$ 1"),
xconn.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)),
corelistener.TrafficLimiterOption(xtraffic.NewTrafficLimiter(
xtraffic.LimitsOption("$$ 1024B 1024B"),
xtraffic.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)),
corelistener.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)
if err := ln.Init(nil); err != nil {
t.Fatalf("init listener: %v", err)
}
defer ln.Close()
client, err := net.Dial("udp", ln.Addr().String())
if err != nil {
t.Fatalf("dial udp listener: %v", err)
}
defer client.Close()
if _, err := client.Write([]byte("packet")); err != nil {
t.Fatalf("write packet: %v", err)
}
conn, err := acceptWithTimeout(t, ln, time.Second)
if err != nil {
t.Fatalf("accept conn: %v", err)
}
defer conn.Close()
packetConn, ok := conn.(net.PacketConn)
if !ok {
t.Fatalf("expected accepted UDP conn with limiters to implement net.PacketConn, got %T", conn)
}
buf := make([]byte, 16)
n, addr, err := packetConn.ReadFrom(buf)
if err != nil {
t.Fatalf("read packet: %v", err)
}
if string(buf[:n]) != "packet" {
t.Fatalf("expected original datagram, got %q", string(buf[:n]))
}
if addr == nil || addr.String() != client.LocalAddr().String() {
t.Fatalf("expected client addr %v, got %v", client.LocalAddr(), addr)
}
}
func TestAcceptAppliesConnLimiterAndReleasesOnClose(t *testing.T) {
ln := NewListener(
corelistener.AddrOption("127.0.0.1:0"),
corelistener.ConnLimiterOption(xconn.NewConnLimiter(
xconn.LimitsOption("$$ 1"),
xconn.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)),
corelistener.LoggerOption(xlogger.NewLogger(xlogger.OutputOption(io.Discard), xlogger.LevelOption(corelogger.ErrorLevel))),
)
if err := ln.Init(nil); err != nil {
t.Fatalf("init listener: %v", err)
}
defer ln.Close()
addr := ln.Addr().String()
client, err := net.Dial("udp", addr)
if err != nil {
t.Fatalf("dial udp listener: %v", err)
}
defer client.Close()
if _, err := client.Write([]byte("first")); err != nil {
t.Fatalf("write first packet: %v", err)
}
first, err := acceptWithTimeout(t, ln, time.Second)
if err != nil {
t.Fatalf("accept first conn: %v", err)
}
blockedClient, err := net.Dial("udp", addr)
if err != nil {
t.Fatalf("dial blocked udp client: %v", err)
}
defer blockedClient.Close()
if _, err := blockedClient.Write([]byte("blocked")); err != nil {
t.Fatalf("write blocked packet: %v", err)
}
blocked, err := acceptWithTimeout(t, ln, time.Second)
if err != nil {
t.Fatalf("expected blocked same-IP pseudo-connection to be returned closed: %v", err)
}
buf := make([]byte, 16)
if _, err := blocked.Read(buf); err == nil {
_ = blocked.Close()
t.Fatalf("expected blocked same-IP pseudo-connection to be closed")
}
packetConn, ok := blocked.(net.PacketConn)
if !ok {
_ = blocked.Close()
t.Fatalf("expected blocked same-IP pseudo-connection to preserve net.PacketConn, got %T", blocked)
}
if _, _, err := packetConn.ReadFrom(buf); err == nil {
_ = blocked.Close()
t.Fatalf("expected blocked same-IP packet connection to be closed")
}
if _, err := packetConn.WriteTo([]byte("blocked"), client.LocalAddr()); err == nil {
_ = blocked.Close()
t.Fatalf("expected blocked same-IP packet write to be closed")
}
_ = blocked.Close()
_ = first.Close()
reopenedClient, err := net.Dial("udp", addr)
if err != nil {
t.Fatalf("dial reopened udp client: %v", err)
}
defer reopenedClient.Close()
if _, err := reopenedClient.Write([]byte("after-close")); err != nil {
t.Fatalf("write after close packet: %v", err)
}
reopened, err := acceptWithTimeout(t, ln, time.Second)
if err != nil {
t.Fatalf("expected same client to be accepted after close: %v", err)
}
_ = reopened.Close()
}
func acceptWithTimeout(t *testing.T, ln corelistener.Listener, timeout time.Duration) (net.Conn, error) {
t.Helper()
type result struct {
conn net.Conn
err error
}
ch := make(chan result, 1)
go func() {
conn, err := ln.Accept()
ch <- result{conn: conn, err: err}
}()
select {
case res := <-ch:
return res.conn, res.err
case <-time.After(timeout):
return nil, net.ErrClosed
}
}
+4
View File
@@ -60,6 +60,8 @@ func GetInt(md metadata.Metadata, keys ...string) (v int) {
}
case int:
v = vv
case float64:
v = int(vv)
case string:
v, _ = strconv.Atoi(vv)
}
@@ -105,6 +107,8 @@ func GetDuration(md metadata.Metadata, keys ...string) (v time.Duration) {
switch vv := md.Get(key).(type) {
case int:
v = time.Duration(vv) * time.Second
case float64:
v = time.Duration(vv) * time.Second
case string:
v, _ = time.ParseDuration(vv)
if v == 0 {
+20 -61
View File
@@ -88,56 +88,45 @@ func (m *GlobalTrafficManager) startReporting() {
// collectAndReport 收集所有服务流量并合并上报
func (m *GlobalTrafficManager) collectAndReport() {
m.mu.Lock()
// 如果没有流量,直接返回
if len(m.serviceTraffic) == 0 {
m.mu.Unlock()
return
}
// 复制当前所有流量数据(避免长时间持锁)
trafficSnapshot := make(map[string]*ServiceTraffic)
reportData := make(map[string]struct {
up int64
down int64
})
reportItems := make([]TrafficReportItem, 0, len(m.serviceTraffic))
for name, traffic := range m.serviceTraffic {
traffic.mu.Lock()
if traffic.UpBytes > 0 || traffic.DownBytes > 0 {
trafficSnapshot[name] = traffic
reportData[name] = struct {
up int64
down int64
}{
up: traffic.UpBytes,
down: traffic.DownBytes,
}
up := traffic.UpBytes
down := traffic.DownBytes
if up > 0 || down > 0 {
traffic.UpBytes = 0
traffic.DownBytes = 0
}
traffic.mu.Unlock()
if up > 0 || down > 0 {
reportItems = append(reportItems, TrafficReportItem{
N: name,
U: up,
D: down,
})
}
}
m.mu.Unlock()
// 如果没有需要上报的流量,返回
if len(reportData) == 0 {
if len(reportItems) == 0 {
return
}
// 构建上报数据数组(保持每个服务独立)
reportItems := make([]TrafficReportItem, 0, len(reportData))
var totalUp, totalDown int64
for serviceName, data := range reportData {
reportItems = append(reportItems, TrafficReportItem{
N: serviceName, // 保持服务名不变
U: data.up,
D: data.down,
})
totalUp += data.up
totalDown += data.down
for _, item := range reportItems {
totalUp += item.U
totalDown += item.D
}
// 批量发送上报请求(一次HTTP请求包含所有服务)
success, err := sendBatchTrafficReport(m.ctx, reportItems)
if err != nil {
fmt.Printf("❌ 全局流量上报失败: %v (总流量: ↑%d ↓%d, %d个服务)\n", err, totalUp, totalDown, len(reportItems))
@@ -146,36 +135,6 @@ func (m *GlobalTrafficManager) collectAndReport() {
if !success {
fmt.Printf("⚠️ 全局流量上报未成功 (总流量: ↑%d ↓%d, %d个服务)\n", totalUp, totalDown, len(reportItems))
return
}
// 上报成功,清空已上报的流量
m.clearReportedTraffic(reportData)
}
// clearReportedTraffic 清空已成功上报的流量
func (m *GlobalTrafficManager) clearReportedTraffic(reportedData map[string]struct {
up int64
down int64
}) {
m.mu.Lock()
defer m.mu.Unlock()
for serviceName, reported := range reportedData {
if traffic, exists := m.serviceTraffic[serviceName]; exists {
traffic.mu.Lock()
// 减去已上报的流量
traffic.UpBytes -= reported.up
traffic.DownBytes -= reported.down
// 如果流量归零,从map中删除该服务记录(避免内存泄漏)
if traffic.UpBytes <= 0 && traffic.DownBytes <= 0 {
traffic.mu.Unlock()
delete(m.serviceTraffic, serviceName)
} else {
traffic.mu.Unlock()
}
}
}
}
+68
View File
@@ -102,3 +102,71 @@ type updateLimiterRequest struct {
type deleteLimiterRequest struct {
Limiter string `json:"limiter"`
}
func createConnLimiter(req createLimiterRequest) error {
name := strings.TrimSpace(req.Data.Name)
if name == "" {
return errors.New("limiter name is required")
}
req.Data.Name = name
if registry.ConnLimiterRegistry().IsRegistered(name) {
return errors.New("conn limiter " + name + " already exists")
}
v := parser.ParseConnLimiter(&req.Data)
if err := registry.ConnLimiterRegistry().Register(name, v); err != nil {
return errors.New("conn limiter " + name + " already exists")
}
if c := config.Global(); c != nil {
c.CLimiters = append(c.CLimiters, &req.Data)
}
return nil
}
func updateConnLimiter(req updateLimiterRequest) error {
name := strings.TrimSpace(req.Limiter)
req.Data.Name = name
if registry.ConnLimiterRegistry().IsRegistered(name) {
registry.ConnLimiterRegistry().Unregister(name)
}
v := parser.ParseConnLimiter(&req.Data)
if err := registry.ConnLimiterRegistry().Register(name, v); err != nil {
return errors.New("conn limiter " + name + " already exists")
}
if c := config.Global(); c != nil {
for i := range c.CLimiters {
if c.CLimiters[i].Name == name {
c.CLimiters[i] = &req.Data
return nil
}
}
c.CLimiters = append(c.CLimiters, &req.Data)
}
return nil
}
func deleteConnLimiter(req deleteLimiterRequest) error {
name := strings.TrimSpace(req.Limiter)
if registry.ConnLimiterRegistry().IsRegistered(name) {
registry.ConnLimiterRegistry().Unregister(name)
}
if c := config.Global(); c != nil {
limiteres := c.CLimiters
c.CLimiters = nil
for _, s := range limiteres {
if s.Name == name {
continue
}
c.CLimiters = append(c.CLimiters, s)
}
}
return nil
}
+237 -8
View File
@@ -145,11 +145,12 @@ type ServiceMonitorCheckResult struct {
}
const (
reporterReadWait = 60 * time.Second
reporterWriteWait = 5 * time.Second
wsPingInterval = 20 * time.Second // 独立 WebSocket ping 间隔
initialBackoff = 2 * time.Second // 重连初始退避
maxBackoff = 2 * time.Minute // 重连最大退避
reporterReadWait = 60 * time.Second
reporterWriteWait = 5 * time.Second
wsPingInterval = 20 * time.Second // 独立 WebSocket ping 间隔
initialBackoff = 2 * time.Second // 重连初始退避
maxBackoff = 2 * time.Minute // 重连最大退避
defaultMetricReportInterval = 5 * time.Second
)
type WebSocketReporter struct {
@@ -189,9 +190,9 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
return &WebSocketReporter{
url: serverURL,
curBackoff: initialBackoff, // 当前退避间隔
pingInterval: 1 * time.Second, // 指标上报间隔(每秒采集)
configInterval: 10 * time.Minute, // 配置上报间隔
curBackoff: initialBackoff, // 当前退避间隔
pingInterval: defaultMetricReportInterval, // 指标上报间隔
configInterval: 10 * time.Minute, // 配置上报间隔
ctx: ctx,
cancel: cancel,
connected: false,
@@ -836,6 +837,18 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
err = w.handleDeleteLimiter(cmd.Data)
response.Type = "DeleteLimitersResponse"
needSaveConfig = true
case "AddCLimiters":
err = w.handleAddCLimiter(cmd.Data)
response.Type = "AddCLimitersResponse"
needSaveConfig = true
case "UpdateCLimiters":
err = w.handleUpdateCLimiter(cmd.Data)
response.Type = "UpdateCLimitersResponse"
needSaveConfig = true
case "DeleteCLimiters":
err = w.handleDeleteCLimiter(cmd.Data)
response.Type = "DeleteCLimitersResponse"
needSaveConfig = true
// TCP Ping 诊断命令(只读,不需要保存配置)
case "TcpPing":
@@ -845,6 +858,13 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
response.Data = tcpPingResult
// needSaveConfig = false (默认值)
// UDP Ping 诊断命令(只读,不需要保存配置)
case "UdpPing":
var udpPingResult TcpPingResponse
udpPingResult, err = w.handleUdpPing(cmd.Data)
response.Type = "UdpPingResponse"
response.Data = udpPingResult
// Service monitor check (read-only)
case "ServiceMonitorCheck":
var checkResult ServiceMonitorCheckResult
@@ -1123,6 +1143,67 @@ func (w *WebSocketReporter) handleDeleteLimiter(data interface{}) error {
return deleteLimiter(deleteReq)
}
func (w *WebSocketReporter) handleAddCLimiter(data interface{}) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("序列化数据失败: %v", err)
}
var limiterConfig config.LimiterConfig
if err := json.Unmarshal(jsonData, &limiterConfig); err != nil {
return fmt.Errorf("解析限流器配置失败: %v", err)
}
req := createLimiterRequest{Data: limiterConfig}
return createConnLimiter(req)
}
func (w *WebSocketReporter) handleUpdateCLimiter(data interface{}) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("序列化数据失败: %v", err)
}
var updateReq struct {
Limiter string `json:"limiter"`
Data config.LimiterConfig `json:"data"`
}
if err := json.Unmarshal(jsonData, &updateReq); err != nil {
var limiterConfig config.LimiterConfig
if err := json.Unmarshal(jsonData, &limiterConfig); err != nil {
return fmt.Errorf("解析更新请求失败: %v", err)
}
updateReq.Limiter = limiterConfig.Name
updateReq.Data = limiterConfig
}
req := updateLimiterRequest{
Limiter: updateReq.Limiter,
Data: updateReq.Data,
}
return updateConnLimiter(req)
}
func (w *WebSocketReporter) handleDeleteCLimiter(data interface{}) error {
jsonData, err := json.Marshal(data)
if err != nil {
return fmt.Errorf("序列化数据失败: %v", err)
}
var deleteReq deleteLimiterRequest
if err := json.Unmarshal(jsonData, &deleteReq); err != nil {
var limiterName string
if err := json.Unmarshal(jsonData, &limiterName); err != nil {
return fmt.Errorf("解析删除请求失败: %v", err)
}
deleteReq.Limiter = limiterName
}
return deleteConnLimiter(deleteReq)
}
// handleSetProtocol 处理设置屏蔽协议的命令
func (w *WebSocketReporter) handleSetProtocol(data interface{}) error {
jsonData, err := json.Marshal(data)
@@ -1610,6 +1691,26 @@ func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls
return reporter
}
var configPersistPath string
// SetConfigPersistPath sets the path where runtime config changes will be
// persisted to disk (gost.json). Called by main during agent startup.
func SetConfigPersistPath(path string) {
configPersistPath = path
config.SetPersistPath(path)
}
// EnableConfigPersist turns on automatic disk persistence after the initial
// config has been loaded and applied.
func EnableConfigPersist() {
config.EnablePersist()
path := config.PersistPath()
if path == "" {
path = configPersistPath
}
fmt.Printf("🔒 节点配置持久化已启用,运行时变更将自动保存到 %s\n", path)
}
// handleTcpPing 处理TCP ping诊断命令
func (w *WebSocketReporter) handleTcpPing(data interface{}) (TcpPingResponse, error) {
jsonData, err := json.Marshal(data)
@@ -1673,6 +1774,64 @@ func (w *WebSocketReporter) handleTcpPing(data interface{}) (TcpPingResponse, er
return response, nil
}
func (w *WebSocketReporter) handleUdpPing(data interface{}) (TcpPingResponse, error) {
jsonData, err := json.Marshal(data)
if err != nil {
return TcpPingResponse{}, fmt.Errorf("序列化UDP ping数据失败: %v", err)
}
var req TcpPingRequest
if err := json.Unmarshal(jsonData, &req); err != nil {
return TcpPingResponse{}, fmt.Errorf("解析UDP ping请求失败: %v", err)
}
if net.ParseIP(req.IP) == nil && !isValidHostname(req.IP) {
return TcpPingResponse{
IP: req.IP,
Port: req.Port,
Success: false,
ErrorMessage: "无效的IP地址或主机名",
RequestId: req.RequestId,
}, nil
}
if req.Port <= 0 || req.Port > 65535 {
return TcpPingResponse{
IP: req.IP,
Port: req.Port,
Success: false,
ErrorMessage: "无效的端口号,范围应为1-65535",
RequestId: req.RequestId,
}, nil
}
if req.Count <= 0 {
req.Count = 4
}
if req.Timeout <= 0 {
req.Timeout = 5000
}
avgTime, packetLoss, err := udpPingHost(req.IP, req.Port, req.Count, req.Timeout)
response := TcpPingResponse{
IP: req.IP,
Port: req.Port,
RequestId: req.RequestId,
}
if err != nil {
response.Success = false
response.ErrorMessage = err.Error()
} else {
response.Success = true
response.AverageTime = avgTime
response.PacketLoss = packetLoss
}
return response, nil
}
// handleServiceMonitorCheck executes a service monitor check on this node.
// It always returns a result (command execution is considered successful even if the check fails).
func (w *WebSocketReporter) handleServiceMonitorCheck(data interface{}) (ServiceMonitorCheckResult, error) {
@@ -1951,6 +2110,76 @@ func tcpPingHost(ip string, port int, count int, timeoutMs int) (float64, float6
return avgTime, packetLoss, nil
}
func udpPingHost(ip string, port int, count int, timeoutMs int) (float64, float64, error) {
var totalTime float64
var successCount int
timeout := time.Duration(timeoutMs) * time.Millisecond
target := net.JoinHostPort(ip, fmt.Sprintf("%d", port))
fmt.Printf("🔍 开始UDP ping测试: %s,次数: %d,超时: %dms\n", target, count, timeoutMs)
if net.ParseIP(ip) == nil {
fmt.Printf("🔍 检测到域名,正在解析DNS...\n")
dnsStart := time.Now()
addrs, err := net.LookupHost(ip)
dnsDuration := time.Since(dnsStart)
if err != nil {
return 0, 100.0, fmt.Errorf("DNS解析失败: %v", err)
}
if len(addrs) == 0 {
return 0, 100.0, fmt.Errorf("DNS解析未返回任何IP地址")
}
fmt.Printf("✅ DNS解析完成 (%.2fms),解析到 %d 个IP: %v\n",
dnsDuration.Seconds()*1000, len(addrs), addrs)
target = net.JoinHostPort(addrs[0], fmt.Sprintf("%d", port))
fmt.Printf("🎯 使用IP地址进行测试: %s\n", target)
} else {
fmt.Printf("🎯 使用IP地址进行测试: %s\n", target)
}
addr, err := net.ResolveUDPAddr("udp", target)
if err != nil {
return 0, 100.0, fmt.Errorf("解析UDP地址失败: %v", err)
}
for i := 0; i < count; i++ {
start := time.Now()
conn, err := net.DialTimeout("udp", addr.String(), timeout)
elapsed := time.Since(start)
if err != nil {
fmt.Printf(" 第%d次UDP连接失败: %v (%.2fms)\n", i+1, err, elapsed.Seconds()*1000)
} else {
fmt.Printf(" 第%d次UDP连接成功: %.2fms\n", i+1, elapsed.Seconds()*1000)
conn.Close()
totalTime += elapsed.Seconds() * 1000
successCount++
}
if i < count-1 {
time.Sleep(100 * time.Millisecond)
}
}
if successCount == 0 {
return 0, 100.0, fmt.Errorf("所有UDP连接尝试都失败")
}
avgTime := totalTime / float64(successCount)
packetLoss := float64(count-successCount) / float64(count) * 100
fmt.Printf("✅ UDP ping完成: 平均连接时间 %.2fms,失败率 %.1f%%\n", avgTime, packetLoss)
return avgTime, packetLoss, nil
}
// isValidHostname 验证主机名格式
func isValidHostname(hostname string) bool {
if len(hostname) == 0 || len(hostname) > 253 {
@@ -110,6 +110,14 @@ func TestSanitizeWebSocketURL(t *testing.T) {
}
}
func TestNewWebSocketReporterUsesReducedMetricInterval(t *testing.T) {
reporter := NewWebSocketReporter("panel.example.com:443", "abc")
if reporter.pingInterval != defaultMetricReportInterval {
t.Fatalf("expected metric interval %s, got %s", defaultMetricReportInterval, reporter.pingInterval)
}
}
func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
err := errors.New("websocket: bad handshake")
resp := &http.Response{
+1 -1
View File
@@ -82,7 +82,7 @@ python with_server.py --backend-port 8080 --frontend-port 5173 -- pytest -v
# Custom server commands
python with_server.py \
--server "make run" --port 6365 --cwd go-backend \
--server "npm run dev" --port 3000 --cwd vite-frontend \
--server "pnpm run dev" --port 3000 --cwd vite-frontend \
-- pytest -v
```
+2 -2
View File
@@ -145,7 +145,7 @@ Examples:
# Custom server configuration
python with_server.py \\
--server "make run" --port 6365 --cwd go-backend \\
--server "npm run dev" --port 3000 --cwd vite-frontend \\
--server "pnpm run dev" --port 3000 --cwd vite-frontend \\
-- pytest -v
# Use custom backend port
@@ -302,7 +302,7 @@ def build_servers(args) -> list[ServerProcess]:
servers.append(
ServerProcess(
command="npm run dev",
command="pnpm run dev",
port=args.frontend_port,
cwd=root / args.frontend_cwd,
env=frontend_env,
+2 -1
View File
@@ -24,7 +24,8 @@ dist-ssr
*.sw?
pnpm-lock.yaml
yarn.lock
package-lock.json
yarn.lock
package-lock.json
bun.lockb
+4 -4
View File
@@ -33,8 +33,8 @@ React dashboard for FLVX. rolldown-vite + TypeScript + Tailwind v4 + shadcn/radi
## Commands
```bash
npm install --legacy-peer-deps
npm run dev # http://0.0.0.0:3000
npm run build # tsc && vite build
npm run lint # eslint --fix
pnpm install
pnpm run dev # http://0.0.0.0:3000
pnpm run build # tsc && vite build
pnpm run lint # eslint --fix
```
+3 -3
View File
@@ -3,11 +3,11 @@ FROM node:20.19.0 AS builder
WORKDIR /app
COPY package*.json ./
RUN npm install --legacy-peer-deps
COPY package.json pnpm-lock.yaml* ./
RUN corepack enable pnpm && pnpm install --frozen-lockfile
COPY . .
RUN npm run build
RUN pnpm run build
# 生产阶段
FROM nginx:stable-alpine AS production-stage
+4 -4
View File
@@ -23,16 +23,16 @@ git clone https://github.com/frontio-ai/vite-template.git
### Install dependencies
You can use one of them `npm`, `yarn`, `pnpm`, `bun`, Example using `npm`:
You can use one of them `npm`, `yarn`, `pnpm`, `bun`, Example using `pnpm`:
```bash
npm install
pnpm install
```
### Run the development server
### Start development server
```bash
npm run dev
pnpm run dev
```
### Setup pnpm (optional)
+2 -1
View File
@@ -52,7 +52,8 @@
"tailwind-merge": "^2.5.5",
"tailwind-variants": "1.0.0",
"tailwindcss": "4.1.11",
"tailwindcss-animate": "^1.0.7"
"tailwindcss-animate": "^1.0.7",
"workbox-window": "^7.4.0"
},
"devDependencies": {
"@eslint/compat": "1.2.8",
+9002
View File
File diff suppressed because it is too large Load Diff
+50 -16
View File
@@ -1,4 +1,10 @@
import { Route, Routes, useLocation, useNavigate, Navigate } from "react-router-dom";
import {
Route,
Routes,
useLocation,
useNavigate,
Navigate,
} from "react-router-dom";
import { useEffect } from "react";
import { AnimatePresence } from "framer-motion";
@@ -21,8 +27,8 @@ import H5SimpleLayout from "@/layouts/h5-simple";
import { isLoggedIn } from "@/utils/auth";
import { siteConfig, updateSiteConfig } from "@/config/site";
import { useH5Mode } from "@/hooks/useH5Mode";
import { SESSION_UPDATED_EVENT } from "@/utils/session";
// 简化的路由保护组件 - 使用 React Router 导航避免循环
const ProtectedRoute = ({
children,
useSimpleLayout = false,
@@ -32,16 +38,8 @@ const ProtectedRoute = ({
useSimpleLayout?: boolean;
skipLayout?: boolean;
}) => {
const authenticated = isLoggedIn();
const isH5 = useH5Mode();
const navigate = useNavigate();
useEffect(() => {
if (!authenticated) {
// 使用 React Router 导航,避免无限跳转
navigate("/", { replace: true });
}
}, [authenticated, navigate]);
const authenticated = isLoggedIn();
if (!authenticated) {
return <Navigate replace to="/" />;
@@ -80,26 +78,60 @@ const LoginRoute = () => {
function App() {
const location = useLocation();
const navigate = useNavigate();
// 全局登录状态监听,当检测到未登录且不在首页时,跳转到首页
useEffect(() => {
const handleSessionUpdate = () => {
if (!isLoggedIn() && location.pathname !== "/") {
navigate("/", { replace: true });
}
};
window.addEventListener(SESSION_UPDATED_EVENT, handleSessionUpdate);
return () => {
window.removeEventListener(SESSION_UPDATED_EVENT, handleSessionUpdate);
};
}, [location.pathname, navigate]);
// 处理自定义背景图片
useEffect(() => {
const updateBg = () => {
const customBg = siteConfig.app_bg_image;
if (customBg) {
if (customBg === "theme") {
document.documentElement.style.removeProperty("--custom-bg-image");
document.documentElement.style.removeProperty("--custom-bg-color");
document.documentElement.classList.add("has-theme-bg");
document.documentElement.classList.remove("has-custom-bg");
} else if (customBg.startsWith("http") || customBg.startsWith("data:") || customBg.startsWith("/") || customBg.startsWith("blob:")) {
document.documentElement.style.setProperty("--custom-bg-image", `url(${customBg})`);
document.documentElement.style.setProperty("--custom-bg-color", "transparent");
} else if (
customBg.startsWith("http") ||
customBg.startsWith("data:") ||
customBg.startsWith("/") ||
customBg.startsWith("blob:")
) {
document.documentElement.style.setProperty(
"--custom-bg-image",
`url(${customBg})`,
);
document.documentElement.style.setProperty(
"--custom-bg-color",
"transparent",
);
document.documentElement.classList.add("has-custom-bg");
document.documentElement.classList.remove("has-theme-bg");
} else {
// Assume solid color like "#ffffff", "white", etc.
document.documentElement.style.setProperty("--custom-bg-image", "none");
document.documentElement.style.setProperty("--custom-bg-color", customBg);
document.documentElement.style.setProperty(
"--custom-bg-image",
"none",
);
document.documentElement.style.setProperty(
"--custom-bg-color",
customBg,
);
document.documentElement.classList.add("has-custom-bg");
document.documentElement.classList.remove("has-theme-bg");
}
@@ -110,8 +142,10 @@ function App() {
document.documentElement.classList.remove("has-theme-bg");
}
};
updateBg();
window.addEventListener("site-config-updated", updateBg);
return () => {
window.removeEventListener("site-config-updated", updateBg);
};
+13 -5
View File
@@ -149,9 +149,12 @@ export const deleteTunnelWithForwards = (data: {
data,
);
export const previewBatchTunnelDelete = (ids: number[]) =>
Network.post<TunnelBatchDeletePreviewApiData>("/tunnel/batch-delete-preview", {
ids,
});
Network.post<TunnelBatchDeletePreviewApiData>(
"/tunnel/batch-delete-preview",
{
ids,
},
);
export const batchDeleteTunnelsWithForwards = (data: {
ids: number[];
action: "replace" | "delete_forwards";
@@ -415,8 +418,11 @@ export interface AnnouncementData {
export const getAnnouncement = () =>
Network.get<AnnouncementData>("/announcement/get");
export const updateAnnouncement = (data: AnnouncementData) =>
Network.post("/announcement/update", data);
export const updateAnnouncement = ({
content,
enabled,
}: Pick<AnnouncementData, "content" | "enabled">) =>
Network.post("/announcement/update", { content, enabled });
export const getNodeMetrics = (
nodeId: number,
@@ -496,12 +502,14 @@ export const getServiceMonitorResults = (
options?: { limit?: number; start?: number; end?: number },
) => {
const params: Record<string, string> = {};
if (options?.start != null && options?.end != null) {
params.start = String(options.start);
params.end = String(options.end);
} else if (options?.limit != null) {
params.limit = String(options.limit);
}
return Network.get<ServiceMonitorResultApiItem[]>(
`/monitor/services/${monitorId}/results`,
params,
+13
View File
@@ -20,6 +20,7 @@ export interface UserApiItem {
num: number;
expTime?: number;
flowResetTime?: number;
maxConn?: number;
inFlow?: number;
outFlow?: number;
dailyQuotaGB?: number;
@@ -70,6 +71,11 @@ export interface ForwardApiItem {
userId?: number;
tunnelId?: number;
speedId?: number | null;
ipMaxConn?: number;
ipSpeedId?: number | null;
ipSpeedLimitName?: string;
maxConn?: number;
proxyProtocol?: number;
inx?: number;
[key: string]: unknown;
}
@@ -200,6 +206,7 @@ export interface UserPackageInfoApiData {
num: number;
expTime?: string;
flowResetTime?: number;
maxConn?: number;
[key: string]: unknown;
};
tunnelPermissions: UserTunnelPermissionApiItem[];
@@ -276,6 +283,7 @@ export interface UserMutationPayload {
num?: number;
expTime?: number | string;
flowResetTime?: number;
maxConn?: number;
dailyQuotaGB?: number;
monthlyQuotaGB?: number;
tunnelFlow?: number;
@@ -338,6 +346,7 @@ export interface UserTunnelAssignPayload {
num?: number;
expTime?: number;
flowResetTime?: number;
maxConn?: number;
status?: number;
speedId?: number | null;
tunnels?: Array<{ tunnelId: number; speedId?: number | null }>;
@@ -366,6 +375,10 @@ export interface ForwardMutationPayload {
remoteAddr?: string;
strategy?: string;
speedId?: number | null;
ipMaxConn?: number;
ipSpeedId?: number | null;
maxConn?: number;
proxyProtocol?: number;
}
export interface SpeedLimitMutationPayload {
File diff suppressed because one or more lines are too long
@@ -7,13 +7,14 @@
* • "Reset to default" option
*/
import type { ThemeMode } from "@/themes/registry";
import React from "react";
import toast from "react-hot-toast";
import { Card, CardBody } from "@/shadcn-bridge/heroui/card";
import { Button } from "@/shadcn-bridge/heroui/button";
import { useThemeContext } from "@/themes/context";
import type { ThemeMode } from "@/themes/registry";
// ─── Constants ──────────────────────────────────────────────────────────────
@@ -39,12 +40,14 @@ export const ThemeSettings: React.FC = () => {
const handleModeChange = (m: ThemeMode) => {
setMode(m);
const label = m === "light" ? "亮色" : m === "dark" ? "暗色" : "跟随系统";
toast.success(`已切换为${label}模式`);
};
const handleThemeSelect = (id: string) => {
switchTheme(id);
const theme = themes.find((t) => t.id === id);
toast.success(`已切换主题「${theme?.name ?? id}」`);
};
@@ -125,7 +128,13 @@ export const ThemeSettings: React.FC = () => {
<div className="flex-1" style={{ background: secondary }} />
<div className="flex-1" style={{ background: success }} />
<div className="flex-1" style={{ background: danger }} />
<div className="flex-1" style={{ background: bg, borderLeft: "1px solid rgba(0,0,0,0.06)" }} />
<div
className="flex-1"
style={{
background: bg,
borderLeft: "1px solid rgba(0,0,0,0.06)",
}}
/>
</div>
{/* Info */}
+8 -9
View File
@@ -1,4 +1,5 @@
import * as React from "react";
import { cn } from "@/lib/utils";
function Card({ className, style, ...props }: React.ComponentProps<"div">) {
@@ -7,16 +8,18 @@ function Card({ className, style, ...props }: React.ComponentProps<"div">) {
className={cn(
"relative flex flex-col transition-all duration-300",
"rounded-[24px] text-card-foreground shadow-[0_12px_40px_rgba(0,0,0,0.12)] dark:shadow-[0_12px_40px_rgba(0,0,0,0.4)]",
className
className,
)}
data-slot="card"
style={{
...style,
background: "linear-gradient(135deg, rgba(255,255,255,0.4) 0%, rgba(255,255,255,0.1) 40%, rgba(255,255,255,0.05) 60%, rgba(255,255,255,0.2) 100%)",
boxShadow: "inset 0 1px 1px rgba(255,255,255,0.8), inset 0 0 0 1px rgba(255,255,255,0.3), inset 0 -1px 1px rgba(0,0,0,0.1), 0 12px 40px rgba(0,0,0,0.12)",
background:
"linear-gradient(135deg, rgba(255,255,255,0.4) 0%, rgba(255,255,255,0.1) 40%, rgba(255,255,255,0.05) 60%, rgba(255,255,255,0.2) 100%)",
boxShadow:
"inset 0 1px 1px rgba(255,255,255,0.8), inset 0 0 0 1px rgba(255,255,255,0.3), inset 0 -1px 1px rgba(0,0,0,0.1), 0 12px 40px rgba(0,0,0,0.12)",
backdropFilter: "blur(24px) saturate(180%)",
WebkitBackdropFilter: "blur(24px) saturate(180%)",
}}
data-slot="card"
{...props}
/>
);
@@ -63,11 +66,7 @@ function CardDescription({ className, ...props }: React.ComponentProps<"p">) {
function CardContent({ className, ...props }: React.ComponentProps<"div">) {
return (
<div
className={cn("p-6", className)}
data-slot="card-content"
{...props}
/>
<div className={cn("p-6", className)} data-slot="card-content" {...props} />
);
}
+5 -3
View File
@@ -62,13 +62,15 @@ function DialogContent({
"rounded-[24px] text-foreground shadow-[0_20px_60px_rgba(0,0,0,0.2)] dark:shadow-[0_20px_60px_rgba(0,0,0,0.5)] border border-white/80 dark:border-white/10",
className,
)}
data-slot="dialog-content"
style={{
background: "linear-gradient(135deg, rgba(255,255,255,0.4) 0%, rgba(255,255,255,0.1) 40%, rgba(255,255,255,0.05) 60%, rgba(255,255,255,0.2) 100%)",
boxShadow: "inset 0 1px 1px rgba(255,255,255,0.8), inset 0 0 0 1px rgba(255,255,255,0.3), inset 0 -1px 1px rgba(0,0,0,0.1), 0 20px 60px rgba(0,0,0,0.2)",
background:
"linear-gradient(135deg, rgba(255,255,255,0.4) 0%, rgba(255,255,255,0.1) 40%, rgba(255,255,255,0.05) 60%, rgba(255,255,255,0.2) 100%)",
boxShadow:
"inset 0 1px 1px rgba(255,255,255,0.8), inset 0 0 0 1px rgba(255,255,255,0.3), inset 0 -1px 1px rgba(0,0,0,0.1), 0 20px 60px rgba(0,0,0,0.2)",
backdropFilter: "blur(24px) saturate(180%)",
WebkitBackdropFilter: "blur(24px) saturate(180%)",
}}
data-slot="dialog-content"
{...props}
>
<div className="flex flex-col flex-1 min-h-0 gap-4 p-6 rounded-[inherit] overflow-hidden bg-white/40 dark:bg-zinc-900/40 w-full h-full relative z-10 pointer-events-auto">
@@ -85,7 +85,7 @@ function DropdownMenuSubContent({
return (
<DropdownMenuPrimitive.SubContent
className={cn(
"z-50 min-w-32 overflow-hidden rounded-md border border-default-200 bg-white p-1 text-foreground shadow-lg data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0 data-[state=closed]:zoom-out-95 data-[state=open]:zoom-in-95 data-[side=top]:slide-in-from-bottom-2 data-[side=bottom]:slide-in-from-top-2 dark:bg-default-50",
"z-50 min-w-32 overflow-y-auto rounded-md border border-default-200 bg-white p-1.5 text-foreground shadow-lg max-h-[--radix-dropdown-menu-content-available-height] [&::-webkit-scrollbar]:hidden [-ms-overflow-style:none] [scrollbar-width:none] data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0 data-[state=closed]:zoom-out-95 data-[state=open]:zoom-in-95 data-[side=top]:slide-in-from-bottom-2 data-[side=bottom]:slide-in-from-top-2 dark:bg-default-50",
className,
)}
data-slot="dropdown-menu-sub-content"
@@ -103,7 +103,7 @@ function DropdownMenuContent({
<DropdownMenuPrimitive.Portal>
<DropdownMenuPrimitive.Content
className={cn(
"z-50 min-w-32 overflow-hidden rounded-md border border-default-200 bg-white p-1 text-foreground shadow-md data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0 data-[state=closed]:zoom-out-95 data-[state=open]:zoom-in-95 data-[side=top]:slide-in-from-bottom-2 data-[side=bottom]:slide-in-from-top-2 dark:bg-default-50",
"z-50 min-w-32 overflow-y-auto rounded-md border border-default-200 bg-white p-1.5 text-foreground shadow-md max-h-[--radix-dropdown-menu-content-available-height] [&::-webkit-scrollbar]:hidden [-ms-overflow-style:none] [scrollbar-width:none] data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0 data-[state=closed]:zoom-out-95 data-[state=open]:zoom-in-95 data-[side=top]:slide-in-from-bottom-2 data-[side=bottom]:slide-in-from-top-2 dark:bg-default-50",
className,
)}
data-slot="dropdown-menu-content"
+4 -2
View File
@@ -32,8 +32,10 @@ const getInitialConfig = () => {
localStorage.getItem(CACHE_PREFIX + "app_favicon") || "";
const cachedAppBgImage =
localStorage.getItem(CACHE_PREFIX + "app_bg_image") || "";
const isCommercial = localStorage.getItem(CACHE_PREFIX + "is_commercial") === "true";
const hideFooterBrand = localStorage.getItem(CACHE_PREFIX + "hide_footer_brand") === "true";
const isCommercial =
localStorage.getItem(CACHE_PREFIX + "is_commercial") === "true";
const hideFooterBrand =
localStorage.getItem(CACHE_PREFIX + "hide_footer_brand") === "true";
if (cachedAppName) {
return {
+12 -9
View File
@@ -198,19 +198,23 @@ export default function AdminLayout({
if (adminFlag) {
setMonitorAllowed(true);
setMonitorAccessReason(null);
return;
}
let cancelled = false;
(async () => {
try {
const res = await getMonitorAccess();
if (cancelled) return;
if (res.code === 0 && res.data) {
setMonitorAllowed(Boolean(res.data.allowed));
setMonitorAccessReason(
res.data.allowed ? null : (res.data.reason || null),
res.data.allowed ? null : res.data.reason || null,
);
return;
}
// Fail open to preserve legacy navigation behavior.
@@ -237,7 +241,6 @@ export default function AdminLayout({
// 退出登录
const handleLogout = () => {
safeLogout();
navigate("/");
};
// 切换移动端菜单
@@ -376,7 +379,7 @@ export default function AdminLayout({
${isMobile ? "fixed h-screen top-0 left-0 rounded-r-3xl" : "relative h-full rounded-3xl"}
${isMobile && !mobileMenuVisible ? "-translate-x-full" : "translate-x-0"}
${isMobile ? "w-64" : isCollapsed ? "w-20" : "w-[260px]"}
bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl
bg-white/70 dark:bg-zinc-900/70 backdrop-blur-3xl
shadow-[0_10px_30px_rgba(0,0,0,0.1)]
border border-white/80 dark:border-white/10
z-50
@@ -398,7 +401,7 @@ export default function AdminLayout({
</div>
{/* 菜单导航 */}
<nav className="flex-1 px-4 overflow-y-auto overflow-x-hidden [scrollbar-width:none]">
<nav className="flex-1 px-4 overflow-y-auto overflow-x-hidden scrollbar-hide">
<ul className="space-y-2">
{filteredMenuItems.map((item) => {
const isActive = location.pathname === item.path;
@@ -408,6 +411,7 @@ export default function AdminLayout({
return (
<li key={item.path}>
<motion.button
aria-disabled={isMonitorBlocked}
className={`
w-full flex items-center p-3 rounded-2xl text-left
relative min-h-[48px] overflow-hidden transition-colors
@@ -415,12 +419,11 @@ export default function AdminLayout({
${
isActive
? "text-primary dark:text-primary-400 font-semibold"
: isMonitorBlocked
? "text-gray-500 dark:text-gray-400 font-medium"
: "text-gray-600 dark:text-gray-300 font-medium"
: isMonitorBlocked
? "text-gray-500 dark:text-gray-400 font-medium"
: "text-gray-600 dark:text-gray-300 font-medium"
}
`}
aria-disabled={isMonitorBlocked}
title={
isCollapsed
? isMonitorBlocked
@@ -551,7 +554,7 @@ export default function AdminLayout({
)}
{/* 主内容 */}
<main className="flex-1 overflow-y-auto [scrollbar-width:none]">
<main className="flex-1 overflow-y-auto scrollbar-hide">
<AnimatePresence mode="wait">
<motion.div
key={location.pathname}
+4 -1
View File
@@ -114,15 +114,18 @@ export default function H5Layout({ children }: { children: React.ReactNode }) {
}
let cancelled = false;
(async () => {
try {
const res = await getMonitorAccess();
if (cancelled) return;
if (res.code === 0 && res.data) {
setMonitorAllowed(Boolean(res.data.allowed));
setMonitorAccessReason(
res.data.allowed ? null : (res.data.reason || null),
res.data.allowed ? null : res.data.reason || null,
);
return;
}
setMonitorAllowed(true);
+4 -2
View File
@@ -13,7 +13,9 @@ const updateSW = registerSW({
toast(
(t) => (
<div className="flex flex-col gap-3">
<span className="text-sm font-medium text-foreground">发现新版本,是否立即刷新以应用更新?</span>
<span className="text-sm font-medium text-foreground">
发现新版本,是否立即刷新以应用更新?
</span>
<div className="flex gap-2 justify-end">
<button
className="px-3 py-1.5 text-xs font-medium bg-primary text-primary-foreground hover:bg-primary/90 rounded-md transition-colors"
@@ -33,7 +35,7 @@ const updateSW = registerSW({
</div>
</div>
),
{ duration: Infinity, position: "bottom-right" }
{ duration: Infinity, position: "bottom-right" },
);
},
});
@@ -1,5 +1,4 @@
import { useState } from "react";
import { useNavigate } from "react-router-dom";
import toast from "react-hot-toast";
import { Button } from "@/shadcn-bridge/heroui/button";
@@ -26,7 +25,6 @@ export default function ChangePasswordPage() {
});
const [loading, setLoading] = useState(false);
const [errors, setErrors] = useState<Partial<PasswordForm>>({});
const navigate = useNavigate();
const validateForm = (): boolean => {
const newErrors: Partial<PasswordForm> = {};
@@ -98,7 +96,6 @@ export default function ChangePasswordPage() {
const logout = () => {
safeLogout();
navigate("/");
};
const handleKeyPress = (e: React.KeyboardEvent) => {
+93 -37
View File
@@ -89,7 +89,8 @@ const CONFIG_ITEMS: ConfigItem[] = [
{
key: "app_bg_image",
label: "自定义背景",
description: "上传自定义背景图片(建议使用深色/浅色均可看清的图片,或使用半透明模糊效果)",
description:
"上传自定义背景图片(建议使用深色/浅色均可看清的图片,或使用半透明模糊效果)",
type: "bg_image",
},
{
@@ -141,7 +142,8 @@ const CONFIG_ITEMS: ConfigItem[] = [
{
key: "monitor_tunnel_quality_enabled",
label: "实时隧道质量检测",
description: "关闭后,前端停止自动刷新,后端停止实时隧道质量探测(全局配置)",
description:
"关闭后,前端停止自动刷新,后端停止实时隧道质量探测(全局配置)",
type: "switch",
},
{
@@ -183,6 +185,13 @@ const CONFIG_ITEMS: ConfigItem[] = [
dependsOn: "github_proxy_enabled",
dependsValue: "true",
},
{
key: "allow_local_remote_addr",
label: "允许转发到本地地址",
description:
"开启后,普通用户创建或编辑规则时可将目标地址指向 127.0.0.1、10.x.x.x、172.16-31.x.x、192.168.x.x 等本地或内网地址。默认关闭以降低开放代理风险。",
type: "switch",
},
];
const BACKUP_TYPE_OPTIONS = [
@@ -217,6 +226,7 @@ const getInitialConfigs = (): Record<string, string> => {
"app_favicon",
"github_proxy_enabled",
"github_proxy_url",
"allow_local_remote_addr",
];
const initialConfigs: Record<string, string> = {};
@@ -388,11 +398,13 @@ export default function ConfigPage() {
const handleActivateLicense = async () => {
if (!licenseKeyInput.trim()) {
toast.error("请输入有效的商业授权码");
return;
}
setActivatingLicense(true);
try {
const res = await activateLicense(licenseKeyInput.trim());
if (res.code === 0) {
toast.success("商业版授权激活成功!");
setLicenseKeyInput("");
@@ -479,7 +491,9 @@ export default function ConfigPage() {
if (changedKeys.includes("monitor_tunnel_quality_enabled")) {
window.dispatchEvent(
new CustomEvent("monitorTunnelQualityEnabledChanged", {
detail: { enabled: configs["monitor_tunnel_quality_enabled"] === "true" },
detail: {
enabled: configs["monitor_tunnel_quality_enabled"] === "true",
},
}),
);
}
@@ -496,7 +510,11 @@ export default function ConfigPage() {
// 检查配置项是否应该显示(依赖检查)
const shouldShowItem = (item: ConfigItem): boolean => {
// 隐藏商业版专属的设置项,它们会在专门的卡片中展示
if (["app_name", "app_logo", "app_favicon", "hide_footer_brand"].includes(item.key)) {
if (
["app_name", "app_logo", "app_favicon", "hide_footer_brand"].includes(
item.key,
)
) {
return false;
}
@@ -555,12 +573,16 @@ export default function ConfigPage() {
}
};
const handleBgImageUpload = async (e: React.ChangeEvent<HTMLInputElement>) => {
const handleBgImageUpload = async (
e: React.ChangeEvent<HTMLInputElement>,
) => {
const file = e.target.files?.[0];
if (!file) return;
if (!file.type.startsWith("image/")) {
toast.error("只能上传图片文件");
return;
}
@@ -568,8 +590,10 @@ export default function ConfigPage() {
try {
const compressedImage = await new Promise<string>((resolve, reject) => {
const reader = new FileReader();
reader.onload = (event) => {
const img = new Image();
img.onload = () => {
const canvas = document.createElement("canvas");
let width = img.width;
@@ -592,14 +616,17 @@ export default function ConfigPage() {
canvas.width = width;
canvas.height = height;
const ctx = canvas.getContext("2d");
if (!ctx) {
reject(new Error("Canvas context is null"));
return;
}
ctx.drawImage(img, 0, 0, width, height);
// Output as webp for better compression
const dataUrl = canvas.toDataURL("image/webp", 0.8);
resolve(dataUrl);
};
img.onerror = () => reject(new Error("图片加载失败"));
@@ -621,62 +648,66 @@ export default function ConfigPage() {
const renderBgImageUploader = () => {
const bgImage = configs["app_bg_image"] || "";
const isImage = bgImage.startsWith("http") || bgImage.startsWith("data:") || bgImage.startsWith("/") || bgImage.startsWith("blob:");
const isImage =
bgImage.startsWith("http") ||
bgImage.startsWith("data:") ||
bgImage.startsWith("/") ||
bgImage.startsWith("blob:");
const isTheme = bgImage === "theme";
const isSolidColor = bgImage && !isImage && !isTheme;
return (
<div className="flex flex-col gap-4 w-full">
<div className="flex flex-wrap items-center gap-4">
<input
type="file"
accept="image/*"
ref={bgImageFileInputRef}
accept="image/*"
className="hidden"
type="file"
onChange={handleBgImageUpload}
/>
<Button
color="primary"
isLoading={bgImageUploading}
variant="flat"
onPress={() => bgImageFileInputRef.current?.click()}
isLoading={bgImageUploading}
>
上传图片
</Button>
<Button
color="secondary"
isDisabled={bgImageUploading || isTheme}
variant="flat"
onPress={() => handleConfigChange("app_bg_image", "theme")}
isDisabled={bgImageUploading || isTheme}
>
自适应纯色 (跟随深色模式)
</Button>
<Button
color="default"
isDisabled={bgImageUploading || bgImage === "#ffffff"}
variant="flat"
onPress={() => handleConfigChange("app_bg_image", "#ffffff")}
isDisabled={bgImageUploading || bgImage === "#ffffff"}
>
白色纯色
</Button>
{bgImage && (
<Button
color="danger"
isDisabled={bgImageUploading}
variant="flat"
onPress={() => handleConfigChange("app_bg_image", "")}
isDisabled={bgImageUploading}
>
恢复默认
</Button>
)}
</div>
{bgImage && isImage && (
<div className="relative rounded-xl overflow-hidden border border-divider">
<img
src={bgImage}
alt="背景预览"
className="w-full max-h-48 object-cover opacity-80"
src={bgImage}
/>
<div className="absolute inset-0 bg-gradient-to-t from-black/50 to-transparent flex items-end p-4">
<span className="text-white text-sm font-medium">预览效果</span>
@@ -686,13 +717,15 @@ export default function ConfigPage() {
{bgImage && isTheme && (
<div className="relative rounded-xl border border-divider bg-background h-32 flex items-center justify-center">
<span className="text-foreground text-sm font-medium">当前为自适应纯色背景</span>
<span className="text-foreground text-sm font-medium">
当前为自适应纯色背景
</span>
</div>
)}
{bgImage && isSolidColor && (
<div
className="relative rounded-xl border border-divider h-32 flex items-center justify-center"
<div
className="relative rounded-xl border border-divider h-32 flex items-center justify-center"
style={{ backgroundColor: bgImage }}
>
<span className="text-gray-500 bg-white/80 dark:bg-black/80 px-2 py-1 rounded text-sm font-medium border border-gray-200 dark:border-gray-800 shadow-sm">
@@ -811,8 +844,8 @@ export default function ConfigPage() {
<div className="flex flex-wrap items-center gap-2">
<Button
color="primary"
isLoading={uploading}
isDisabled={isCommercialDisabled}
isLoading={uploading}
size="sm"
variant="flat"
onPress={() => triggerBrandFilePicker(key)}
@@ -853,7 +886,10 @@ export default function ConfigPage() {
const renderConfigItem = (item: ConfigItem) => {
const isChanged =
hasChanges && configs[item.key] !== originalConfigs[item.key];
const isCommercialDisabled = ["app_name", "app_logo", "app_favicon", "hide_footer_brand"].includes(item.key) && configs.is_commercial !== "true";
const isCommercialDisabled =
["app_name", "app_logo", "app_favicon", "hide_footer_brand"].includes(
item.key,
) && configs.is_commercial !== "true";
switch (item.type) {
case "bg_image":
@@ -872,13 +908,15 @@ export default function ConfigPage() {
? "border-warning-300 data-[hover=true]:border-warning-400"
: "",
}}
description={
isCommercialDisabled ? "需商业版授权才能修改此项" : undefined
}
isDisabled={isCommercialDisabled}
placeholder={item.placeholder}
size="md"
value={configs[item.key] || ""}
variant="bordered"
onChange={(e) => handleConfigChange(item.key, e.target.value)}
isDisabled={isCommercialDisabled}
description={isCommercialDisabled ? "需商业版授权才能修改此项" : undefined}
/>
);
@@ -888,12 +926,12 @@ export default function ConfigPage() {
classNames={{
wrapper: isChanged ? "border-warning-300" : "",
}}
isDisabled={isCommercialDisabled}
isSelected={configs[item.key] === "true"}
size="sm"
onValueChange={(checked) =>
handleConfigChange(item.key, checked ? "true" : "false")
}
isDisabled={isCommercialDisabled}
>
<span className="text-sm text-gray-700 dark:text-gray-300">
{configs[item.key] === "true" ? "已启用" : "已禁用"}
@@ -1126,7 +1164,8 @@ export default function ConfigPage() {
<div>
<h2 className="text-xl font-semibold">商业版授权</h2>
<p className="text-sm text-gray-600 dark:text-gray-400">
激活商业版授权以解锁自定义品牌功能(替换 Logo、应用名称,移除底部版权信息等)
激活商业版授权以解锁自定义品牌功能(替换
Logo、应用名称,移除底部版权信息等)
</p>
</div>
</div>
@@ -1135,30 +1174,47 @@ export default function ConfigPage() {
<CardBody className="pt-8">
<div className="flex items-end gap-3 max-w-lg mb-6">
<Input
description={
configs.is_commercial === "true"
? configs.license_expiry && configs.license_expiry !== "never"
? `已激活商业版授权 (有效期至: ${new Date(configs.license_expiry).toLocaleDateString()})`
: "已激活商业版授权 (永久有效)"
: "需商业授权才能修改站名、图标并隐藏页脚品牌"
}
isDisabled={configs.is_commercial === "true"}
label="授权激活码"
placeholder="请输入 FLVX- 开头的商业授权码"
value={licenseKeyInput}
variant="bordered"
onChange={(e) => setLicenseKeyInput(e.target.value)}
isDisabled={configs.is_commercial === "true"}
description={configs.is_commercial === "true" ? (configs.license_expiry && configs.license_expiry !== "never" ? `已激活商业版授权 (有效期至: ${new Date(configs.license_expiry).toLocaleDateString()})` : "已激活商业版授权 (永久有效)") : "需商业授权才能修改站名、图标并隐藏页脚品牌"}
/>
<Button
color="primary"
className="mb-6"
isDisabled={configs.is_commercial === "true" || !licenseKeyInput.trim()}
color="primary"
isDisabled={
configs.is_commercial === "true" || !licenseKeyInput.trim()
}
isLoading={activatingLicense}
onPress={handleActivateLicense}
>
{configs.is_commercial === "true" ? "已授权" : "激活授权"}
</Button>
</div>
{configs.is_commercial === "true" && (
<div className="space-y-6">
<Divider className="my-2" />
<h3 className="text-md font-medium text-default-700">白标与品牌设置</h3>
{CONFIG_ITEMS.filter(item => ["app_name", "app_logo", "app_favicon", "hide_footer_brand"].includes(item.key)).map((item, index) => (
<h3 className="text-md font-medium text-default-700">
白标与品牌设置
</h3>
{CONFIG_ITEMS.filter((item) =>
[
"app_name",
"app_logo",
"app_favicon",
"hide_footer_brand",
].includes(item.key),
).map((item, index) => (
<div key={item.key}>
<div className="grid grid-cols-1 gap-6 md:grid-cols-[1fr_2fr]">
<div className="space-y-1">
@@ -1511,18 +1567,18 @@ export default function ConfigPage() {
<AnimatePresence>
{hasChanges && (
<motion.div
initial={{ y: 100, opacity: 0 }}
animate={{ y: 0, opacity: 1 }}
exit={{ y: 100, opacity: 0 }}
transition={{ type: "spring", damping: 20, stiffness: 300 }}
className="fixed bottom-6 right-6 z-50"
exit={{ y: 100, opacity: 0 }}
initial={{ y: 100, opacity: 0 }}
transition={{ type: "spring", damping: 20, stiffness: 300 }}
>
<Button
isIconOnly
color="primary"
size="lg"
className="w-12 h-12 rounded-full shadow-lg"
color="primary"
isLoading={saving}
size="lg"
onPress={handleSave}
>
{!saving && <SaveIcon className="w-5 h-5" />}
+36 -8
View File
@@ -50,7 +50,6 @@ export default function DashboardPage() {
const handleLogout = () => {
safeLogout();
toast.success("已退出登录");
navigate("/");
};
const {
@@ -651,25 +650,44 @@ export default function DashboardPage() {
{/* 顶部个人状态和下拉菜单 */}
<div className="flex justify-between items-center mb-6">
<div>
<h1 className="text-2xl lg:text-3xl font-bold text-foreground">Good morning, {username}</h1>
<p className="text-sm text-default-500 mt-1">Here's what's happening with your network today.</p>
<h1 className="text-2xl lg:text-3xl font-bold text-foreground">
Good morning, {username}
</h1>
<p className="text-sm text-default-500 mt-1">
Here&apos;s what&apos;s happening with your network today.
</p>
</div>
<div className="flex items-center gap-3">
<Dropdown placement="bottom-end">
<DropdownTrigger>
<Button
className="rounded-full w-10 h-10 min-w-0 bg-primary text-white font-bold text-sm shadow-[0_4px_12px_rgba(0,122,255,0.3)]"
isIconOnly
className="rounded-full w-10 h-10 min-w-0 bg-primary text-white font-bold text-sm shadow-[0_4px_12px_rgba(0,122,255,0.3)]"
variant="solid"
>
{username.slice(0, 2).toUpperCase()}
</Button>
</DropdownTrigger>
<DropdownMenu aria-label="用户菜单" className="bg-white/80 dark:bg-zinc-900/80 backdrop-blur-xl border border-white/50 dark:border-white/10 rounded-2xl shadow-[0_10px_30px_rgba(0,0,0,0.1)]">
<DropdownMenu
aria-label="用户菜单"
className="bg-white/80 dark:bg-zinc-900/80 backdrop-blur-xl border border-white/50 dark:border-white/10 rounded-2xl shadow-[0_10px_30px_rgba(0,0,0,0.1)]"
>
<DropdownItem
key="profile"
startContent={
<svg className="w-4 h-4" fill="none" stroke="currentColor" viewBox="0 0 24 24"><path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M16 7a4 4 0 11-8 0 4 4 0 018 0zM12 14a7 7 0 00-7 7h14a7 7 0 00-7-7z" /></svg>
<svg
className="w-4 h-4"
fill="none"
stroke="currentColor"
viewBox="0 0 24 24"
>
<path
d="M16 7a4 4 0 11-8 0 4 4 0 018 0zM12 14a7 7 0 00-7 7h14a7 7 0 00-7-7z"
strokeLinecap="round"
strokeLinejoin="round"
strokeWidth={2}
/>
</svg>
}
onPress={() => navigate("/profile")}
>
@@ -680,7 +698,17 @@ export default function DashboardPage() {
className="text-danger"
color="danger"
startContent={
<svg className="w-4 h-4" fill="currentColor" viewBox="0 0 20 20"><path fillRule="evenodd" d="M3 3a1 1 0 00-1 1v12a1 1 0 102 0V4a1 1 0 00-1-1zm10.293 9.293a1 1 0 001.414 1.414l3-3a1 1 0 000-1.414l-3-3a1 1 0 10-1.414 1.414L14.586 9H7a1 1 0 100 2h7.586l-1.293 1.293z" clipRule="evenodd" /></svg>
<svg
className="w-4 h-4"
fill="currentColor"
viewBox="0 0 20 20"
>
<path
clipRule="evenodd"
d="M3 3a1 1 0 00-1 1v12a1 1 0 102 0V4a1 1 0 00-1-1zm10.293 9.293a1 1 0 001.414 1.414l3-3a1 1 0 000-1.414l-3-3a1 1 0 10-1.414 1.414L14.586 9H7a1 1 0 100 2h7.586l-1.293 1.293z"
fillRule="evenodd"
/>
</svg>
}
onPress={handleLogout}
>
@@ -1176,7 +1204,7 @@ export default function DashboardPage() {
</Button>
</div>
<div className="space-y-2 max-h-60 overflow-y-auto">
<div className="space-y-2 max-h-60 overflow-y-auto scrollbar-hide">
{addressList.map((item) => (
<div
key={item.id}
@@ -1,4 +1,8 @@
import type { AnnouncementData } from "@/api";
import ReactMarkdown from "react-markdown";
import remarkGfm from "remark-gfm";
import { Button } from "@/shadcn-bridge/heroui/button";
import {
Modal,
@@ -7,8 +11,6 @@ import {
ModalFooter,
ModalHeader,
} from "@/shadcn-bridge/heroui/modal";
import ReactMarkdown from "react-markdown";
import remarkGfm from "remark-gfm";
interface AnnouncementModalProps {
announcement: AnnouncementData;
@@ -24,7 +26,11 @@ export const AnnouncementModal = ({
onDontShowAgain,
}: AnnouncementModalProps) => {
return (
<Modal isOpen={isOpen} onOpenChange={(open) => !open && onClose()} size="2xl">
<Modal
isOpen={isOpen}
size="2xl"
onOpenChange={(open) => !open && onClose()}
>
<ModalContent>
<ModalHeader className="flex flex-col gap-1">平台公告</ModalHeader>
<ModalBody>
@@ -50,7 +50,10 @@ export const FlowChartCard = ({
) : (
<div className="h-64 lg:h-80 w-full">
<ResponsiveContainer height="100%" width="100%">
<LineChart data={chartData} margin={{ top: 20, right: 0, left: -20, bottom: 0 }}>
<LineChart
data={chartData}
margin={{ top: 20, right: 0, left: -20, bottom: 0 }}
>
<XAxis
axisLine={false}
dataKey="time"
@@ -106,11 +109,7 @@ export const FlowChartCard = ({
}}
cursor={{ fill: "rgba(0, 122, 255, 0.05)" }}
/>
<Line
dataKey="flow"
stroke="#007aff"
type="monotone"
/>
<Line dataKey="flow" stroke="#007aff" type="monotone" />
</LineChart>
</ResponsiveContainer>
</div>
@@ -25,9 +25,7 @@ export const MetricCard = ({
<p className="text-sm font-medium text-default-600 truncate">
{title}
</p>
<div
className={`p-2 rounded-xl flex-shrink-0 ${iconClassName}`}
>
<div className={`p-2 rounded-xl flex-shrink-0 ${iconClassName}`}>
{icon}
</div>
</div>
@@ -35,9 +33,7 @@ export const MetricCard = ({
{value}
</p>
</div>
<div className="mt-auto pt-4">
{bottomContent}
</div>
<div className="mt-auto pt-4">{bottomContent}</div>
</CardBody>
</Card>
);
@@ -349,17 +349,18 @@ export const useDashboardData = (): DashboardDataState => {
if (res.code === 0 && res.data && res.data.enabled === 1) {
setAnnouncement(res.data);
try {
const storedTimeStr = localStorage.getItem("flvx_announcement_seen_time");
const storedTimeStr = localStorage.getItem(
"flvx_announcement_seen_time",
);
const storedTime = storedTimeStr ? parseInt(storedTimeStr, 10) : 0;
const updateTime = res.data.update_time || 0;
if (updateTime > storedTime) {
setIsAnnouncementModalOpen(true);
}
} catch (err) {
console.warn("Failed to read localStorage for announcement state", err);
} catch {
setIsAnnouncementModalOpen(true);
}
} else {
@@ -376,9 +377,12 @@ export const useDashboardData = (): DashboardDataState => {
setIsAnnouncementModalOpen(false);
if (announcement && announcement.update_time) {
try {
localStorage.setItem("flvx_announcement_seen_time", announcement.update_time.toString());
} catch (err) {
console.warn("Failed to set localStorage for announcement state", err);
localStorage.setItem(
"flvx_announcement_seen_time",
announcement.update_time.toString(),
);
} catch {
// Ignore localStorage write failures and keep the modal dismissed.
}
}
}, [announcement]);
+214 -50
View File
@@ -54,6 +54,7 @@ import { Switch } from "@/shadcn-bridge/heroui/switch";
import { Alert } from "@/shadcn-bridge/heroui/alert";
import { Progress } from "@/shadcn-bridge/heroui/progress";
import { Checkbox } from "@/shadcn-bridge/heroui/checkbox";
import { Accordion, AccordionItem } from "@/shadcn-bridge/heroui/accordion";
import {
createForward,
getForwardList,
@@ -118,11 +119,16 @@ interface Forward {
outFlow: number;
serviceRunning: boolean;
federationShareFlow?: number;
maxConn?: number;
createdTime: string;
userName?: string;
userId?: number;
inx?: number;
speedId?: number | null;
ipMaxConn?: number;
ipSpeedId?: number | null;
ipSpeedLimitName?: string;
proxyProtocol?: number;
}
interface Tunnel {
@@ -157,6 +163,10 @@ interface ForwardForm {
interfaceName?: string;
strategy: string;
speedId: number | null;
ipMaxConn?: number;
ipSpeedId: number | null;
maxConn?: number;
proxyProtocol?: number;
}
interface ForwardUserGroup {
@@ -573,6 +583,21 @@ const mapForwardApiItems = (items: ForwardApiItem[]): Forward[] => {
typeof forward.speedId === "number" || forward.speedId === null
? forward.speedId
: undefined,
ipMaxConn:
typeof forward.ipMaxConn === "number" ? forward.ipMaxConn : undefined,
ipSpeedId:
typeof forward.ipSpeedId === "number" || forward.ipSpeedId === null
? forward.ipSpeedId
: undefined,
ipSpeedLimitName:
typeof forward.ipSpeedLimitName === "string"
? forward.ipSpeedLimitName
: undefined,
maxConn: typeof forward.maxConn === "number" ? forward.maxConn : undefined,
proxyProtocol:
typeof forward.proxyProtocol === "number"
? forward.proxyProtocol
: undefined,
serviceRunning: forward.status === 1,
}));
};
@@ -1306,6 +1331,10 @@ export default function ForwardPage() {
interfaceName: "",
strategy: "fifo",
speedId: null,
ipMaxConn: 0,
ipSpeedId: null,
maxConn: 0,
proxyProtocol: 0,
});
const [inIpTouched, setInIpTouched] = useState(false);
@@ -2018,6 +2047,7 @@ export default function ForwardPage() {
};
const selectedSpeedId = normalizeSpeedId(form.speedId);
const selectedIPSpeedId = normalizeSpeedId(form.ipSpeedId);
const validateForm = (): boolean => {
const newErrors: { [key: string]: string } = {};
@@ -2093,6 +2123,9 @@ export default function ForwardPage() {
interfaceName: "",
strategy: "fifo",
speedId: null,
ipMaxConn: 0,
ipSpeedId: null,
proxyProtocol: 0,
});
setErrors({});
setModalOpen(true);
@@ -2113,6 +2146,10 @@ export default function ForwardPage() {
interfaceName: forward.interfaceName || "",
strategy: forward.strategy || "fifo",
speedId: normalizeSpeedId(forward.speedId),
ipMaxConn: forward.ipMaxConn ?? 0,
ipSpeedId: normalizeSpeedId(forward.ipSpeedId),
maxConn: forward.maxConn ?? 0,
proxyProtocol: forward.proxyProtocol ?? 0,
});
setErrors({});
setModalOpen(true);
@@ -2230,6 +2267,8 @@ export default function ForwardPage() {
let res: { code: number; msg: string };
const normalizedSpeedId = normalizeSpeedId(form.speedId);
const speedLimitAutoCleared = isMissingSpeedLimit(form.speedId);
const normalizedIPSpeedId = normalizeSpeedId(form.ipSpeedId);
const ipSpeedLimitAutoCleared = isMissingSpeedLimit(form.ipSpeedId);
if (isEdit) {
const updateData = {
@@ -2241,6 +2280,10 @@ export default function ForwardPage() {
remoteAddr: processedRemoteAddr,
strategy: addressCount > 1 ? form.strategy : "fifo",
speedId: normalizedSpeedId,
ipMaxConn: form.ipMaxConn,
...(isAdmin ? { ipSpeedId: normalizedIPSpeedId } : {}),
maxConn: form.maxConn,
proxyProtocol: form.proxyProtocol,
};
res = await updateForward(updateData);
@@ -2253,11 +2296,14 @@ export default function ForwardPage() {
remoteAddr: processedRemoteAddr,
strategy: addressCount > 1 ? form.strategy : "fifo",
speedId: normalizedSpeedId,
ipMaxConn: form.ipMaxConn,
...(isAdmin ? { ipSpeedId: normalizedIPSpeedId } : {}),
maxConn: form.maxConn,
proxyProtocol: form.proxyProtocol,
};
res = await createForward(createData);
}
if (res.code === 0) {
const warningItems = Array.isArray((res as any).data?.warnings)
? (res as any).data.warnings
@@ -2279,6 +2325,12 @@ export default function ForwardPage() {
duration: 5000,
});
}
if (isAdmin && ipSpeedLimitAutoCleared) {
toast("所选每 IP 限速规则不存在,已自动清除为不限速", {
icon: "⚠️",
duration: 5000,
});
}
toast.success(isEdit ? "修改成功" : "创建成功");
setModalOpen(false);
await refreshForwardList(false);
@@ -2582,7 +2634,7 @@ export default function ForwardPage() {
try {
document.execCommand("copy");
toast.success(`已复制${label}`);
} catch (err) {
} catch {
toast.error("复制失败");
}
document.body.removeChild(textArea);
@@ -4176,7 +4228,7 @@ export default function ForwardPage() {
{sortedForwards.length} 条规则
</span>
</div>
<Card className="overflow-hidden rounded-2xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_15px_35px_rgba(0,0,0,0.1)]">
<Card className="rounded-2xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_15px_35px_rgba(0,0,0,0.1)]">
<DndContext
collisionDetection={pointerWithin}
sensors={sensors}
@@ -4188,8 +4240,10 @@ export default function ForwardPage() {
>
<Table
aria-label="全部规则列表"
className="table-fixed min-w-[1160px]"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-[24px]",
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-[24px]",
th: "bg-default-100/50 text-default-600 font-semibold text-sm border-b border-divider py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/10 dark:hover:bg-white/5 transition-colors group",
@@ -4327,7 +4381,7 @@ export default function ForwardPage() {
return (
<div
key={`grouped-table-${group.userId}-${group.userName}`}
className="overflow-hidden rounded-2xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_15px_35px_rgba(0,0,0,0.1)]"
className="rounded-2xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_15px_35px_rgba(0,0,0,0.1)]"
>
<div className="flex items-center justify-between border-b border-white/20 dark:border-white/10 bg-white/20 dark:bg-black/20 backdrop-blur-3xl px-5 py-4">
<div className="flex items-center gap-2">
@@ -4418,7 +4472,7 @@ export default function ForwardPage() {
headerClassName="flex items-center justify-between border-b border-white/20 dark:border-white/10 bg-white/20 dark:bg-black/20 backdrop-blur-3xl px-4 py-2.5"
titleClassName="truncate text-sm font-semibold text-default-700"
tunnel={tunnel}
wrapperClassName="overflow-hidden rounded-2xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_15px_35px_rgba(0,0,0,0.1)]"
wrapperClassName="rounded-2xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_15px_35px_rgba(0,0,0,0.1)]"
onToggleCollapsed={() =>
toggleTunnelGroupCollapsed(
group.userId,
@@ -4439,7 +4493,8 @@ export default function ForwardPage() {
aria-label={`${group.userName}-${tunnel.tunnelName}规则列表`}
className={`table-fixed ${FORWARD_GROUPED_TABLE_MIN_WIDTH_CLASS}`}
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-2xl",
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
@@ -4732,38 +4787,6 @@ export default function ForwardPage() {
}
/>
{isAdmin && (
<Select
label="规则限速"
placeholder="不限速"
selectedKeys={
selectedSpeedId !== null
? [selectedSpeedId.toString()]
: []
}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as
| string
| undefined;
setForm((prev) => ({
...prev,
speedId: selectedKey ? Number(selectedKey) : null,
}));
}}
>
{availableSpeedLimits.map((speedLimit) => (
<SelectItem
key={speedLimit.id.toString()}
textValue={speedLimit.name}
>
{speedLimit.name}
</SelectItem>
))}
</Select>
)}
<Select
description={
isEdit
@@ -4891,6 +4914,147 @@ export default function ForwardPage() {
<SelectItem key="hash">哈希模式 - IP哈希</SelectItem>
</Select>
)}
<Accordion className="px-0" variant="light">
<AccordionItem
key="advanced"
aria-label="高级设置"
className="border-b-0 [&_[data-slot=accordion-trigger]]:no-underline [&_[data-slot=accordion-trigger]]:hover:no-underline"
title={
<span className="text-small text-default-500 font-medium">
高级设置
</span>
}
>
<div className="space-y-4 pb-2">
<Input
description="大于 0 时优先于用户全局限制;0 或空表示使用用户全局限制,用户也为 0 时不限制。"
label="最大连接数"
min="0"
placeholder="0 或空表示使用用户全局限制"
type="number"
value={
form.maxConn === 0 ? "" : String(form.maxConn || "")
}
variant="bordered"
onChange={(e) => {
const value = Math.max(
Number(e.target.value) || 0,
0,
);
setForm((prev) => ({ ...prev, maxConn: value }));
}}
/>
<Input
description="每个客户端 IP 可同时建立的最大连接数;0 或空表示不限制。"
label="每 IP 最大连接数"
min="0"
placeholder="0 或空表示不限制"
type="number"
value={
form.ipMaxConn === 0
? ""
: String(form.ipMaxConn || "")
}
variant="bordered"
onChange={(e) => {
const value = Math.max(
Number(e.target.value) || 0,
0,
);
setForm((prev) => ({ ...prev, ipMaxConn: value }));
}}
/>
<Select
description="启用 PROXY protocol,用于透传客户端真实 IP"
label="Proxy Protocol"
placeholder="禁用"
selectedKeys={[String(form.proxyProtocol || 0)]}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as string;
setForm((prev) => ({
...prev,
proxyProtocol: Number(selectedKey),
}));
}}
>
<SelectItem key="0">禁用</SelectItem>
<SelectItem key="1">Version 1</SelectItem>
<SelectItem key="2">Version 2</SelectItem>
</Select>
{isAdmin && (
<Select
label="规则限速"
placeholder="不限速"
selectedKeys={
selectedSpeedId !== null
? [selectedSpeedId.toString()]
: []
}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as
| string
| undefined;
setForm((prev) => ({
...prev,
speedId: selectedKey
? Number(selectedKey)
: null,
}));
}}
>
{availableSpeedLimits.map((speedLimit) => (
<SelectItem
key={speedLimit.id.toString()}
textValue={speedLimit.name}
>
{speedLimit.name}
</SelectItem>
))}
</Select>
)}
{isAdmin && (
<Select
description="每个客户端 IP 独享该限速规则;不选择表示不限制。"
label="每 IP 限速"
placeholder="不限速"
selectedKeys={
selectedIPSpeedId !== null
? [selectedIPSpeedId.toString()]
: []
}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as
| string
| undefined;
setForm((prev) => ({
...prev,
ipSpeedId: selectedKey
? Number(selectedKey)
: null,
}));
}}
>
{availableSpeedLimits.map((speedLimit) => (
<SelectItem
key={speedLimit.id.toString()}
textValue={speedLimit.name}
>
{speedLimit.name}
</SelectItem>
))}
</Select>
)}
</div>
</AccordionItem>
</Accordion>
</div>
</ModalBody>
<ModalFooter>
@@ -4976,7 +5140,7 @@ export default function ForwardPage() {
</Button>
</div>
<div className="space-y-2 max-h-60 overflow-y-auto">
<div className="space-y-2 max-h-60 overflow-y-auto scrollbar-hide">
{addressList.map((item) => (
<div
key={item.id}
@@ -5281,7 +5445,7 @@ export default function ForwardPage() {
</div>
<div
className="max-h-40 overflow-y-auto space-y-1"
className="max-h-40 overflow-y-auto space-y-1 scrollbar-hide"
style={{
scrollbarWidth: "thin",
scrollbarColor: "rgb(156 163 175) transparent",
@@ -5380,7 +5544,7 @@ export default function ForwardPage() {
<Modal
backdrop="blur"
classNames={{
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden [&>div]:bg-content1 [&>div]:dark:bg-content1",
}}
isOpen={diagnosisModalOpen}
placement="center"
@@ -5398,7 +5562,7 @@ export default function ForwardPage() {
<ModalContent>
{(onClose) => (
<>
<ModalHeader className="flex flex-col gap-1 bg-content1 border-b border-divider">
<ModalHeader className="flex flex-col gap-1 border-b border-divider">
<h2 className="text-xl font-bold">规则诊断结果</h2>
{currentDiagnosisForward && (
<div className="flex items-center gap-2 min-w-0">
@@ -5447,7 +5611,7 @@ export default function ForwardPage() {
{/* 统计摘要 */}
<div className="grid grid-cols-3 gap-3">
<div className="text-center p-3 bg-default-100 dark:bg-gray-800 rounded-lg border border-divider">
<div className="text-center p-3 bg-default-100 rounded-lg border border-divider">
<div className="text-2xl font-bold text-foreground">
{diagnosisProgress.total > 0
? diagnosisProgress.total
@@ -5519,7 +5683,7 @@ export default function ForwardPage() {
return (
<div
key={title}
className="border border-divider rounded-lg overflow-hidden bg-white dark:bg-gray-800"
className="border border-divider rounded-lg overflow-hidden"
>
<div className="bg-primary/10 dark:bg-primary/20 px-3 py-2 border-b border-divider">
<h3 className="text-sm font-semibold text-primary">
@@ -5527,7 +5691,7 @@ export default function ForwardPage() {
</h3>
</div>
<table className="w-full text-sm">
<thead className="bg-default-100 dark:bg-gray-700">
<thead className="bg-default-100">
<tr>
<th className="px-3 py-2 text-left font-semibold text-xs">
路径
@@ -5546,7 +5710,7 @@ export default function ForwardPage() {
</th>
</tr>
</thead>
<tbody className="divide-y divide-divider bg-white dark:bg-gray-800">
<tbody className="divide-y divide-divider bg-content1">
{results.map((result, index) => {
const isDiagnosing = Boolean(
result.diagnosing,
@@ -5565,7 +5729,7 @@ export default function ForwardPage() {
isDiagnosing
? "bg-warning-50 dark:bg-warning-900/20"
: isSuccess
? "bg-white dark:bg-gray-800"
? "bg-content1"
: "bg-danger-50 dark:bg-danger-900/30"
}`}
>
@@ -5752,7 +5916,7 @@ export default function ForwardPage() {
isDiagnosing
? "border-warning-200 dark:border-warning-300/30 bg-warning-50 dark:bg-warning-900/20"
: isSuccess
? "border-divider bg-white dark:bg-gray-800"
? "border-divider bg-content1"
: "border-danger-200 dark:border-danger-300/30 bg-danger-50 dark:bg-danger-900/30"
}`}
>
@@ -5940,7 +6104,7 @@ export default function ForwardPage() {
</div>
)}
</ModalBody>
<ModalFooter className="bg-transparent border-t border-white/20 dark:border-white/10">
<ModalFooter className="border-t border-divider">
<Button variant="light" onPress={onClose}>
关闭
</Button>
+435 -433
View File
@@ -24,7 +24,6 @@ import {
TableRow,
} from "@/shadcn-bridge/heroui/table";
import { Chip } from "@/shadcn-bridge/heroui/chip";
import {
assignGroupPermission,
assignTunnelsToGroup,
@@ -477,451 +476,454 @@ export default function GroupPage() {
) : (
<>
<Card>
<CardHeader className="flex flex-row items-center gap-3 pb-2">
<h3 className="text-lg font-semibold">隧道分组</h3>
<Button
className="h-7 px-3 text-xs font-medium min-w-0 shadow-sm"
color="primary"
size="sm"
onPress={openCreateTunnelGroup}
>
新建
</Button>
</CardHeader>
<CardBody>
<Table
aria-label="隧道分组列表"
<CardHeader className="flex flex-row items-center gap-3 pb-2">
<h3 className="text-lg font-semibold">隧道分组</h3>
<Button
className="h-7 px-3 text-xs font-medium min-w-0 shadow-sm"
color="primary"
size="sm"
onPress={openCreateTunnelGroup}
>
新建
</Button>
</CardHeader>
<CardBody>
<Table
aria-label="隧道分组列表"
classNames={{
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
}}
>
<TableHeader>
<TableColumn>名称</TableColumn>
<TableColumn>隧道</TableColumn>
<TableColumn>状态</TableColumn>
<TableColumn>创建时间</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody emptyContent="暂无隧道分组" items={tunnelGroups}>
{(item) => (
<TableRow key={item.id}>
<TableCell>{item.name}</TableCell>
<TableCell>
{item.tunnelNames.length > 0
? item.tunnelNames.join("、")
: "-"}
</TableCell>
<TableCell>
<Chip
color={item.status === 1 ? "success" : "danger"}
size="sm"
>
{item.status === 1 ? "启用" : "停用"}
</Chip>
</TableCell>
<TableCell>{formatDate(item.createdTime)}</TableCell>
<TableCell>
<div className="flex gap-2">
<Button
size="sm"
variant="flat"
onPress={() => openAssignTunnels(item)}
>
分配隧道
</Button>
<Button
size="sm"
variant="light"
onPress={() => openEditTunnelGroup(item)}
>
编辑
</Button>
<Button
color="danger"
size="sm"
variant="light"
onPress={() => handleDeleteTunnelGroup(item.id)}
>
删除
</Button>
</div>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</CardBody>
</Card>
<Card>
<CardHeader className="flex flex-row items-center gap-3 pb-2">
<h3 className="text-lg font-semibold">用户分组</h3>
<Button
className="h-7 px-3 text-xs font-medium min-w-0 shadow-sm"
color="primary"
size="sm"
onPress={openCreateUserGroup}
>
新建
</Button>
</CardHeader>
<CardBody>
<Table
aria-label="用户分组列表"
classNames={{
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
}}
>
<TableHeader>
<TableColumn>名称</TableColumn>
<TableColumn>用户</TableColumn>
<TableColumn>状态</TableColumn>
<TableColumn>创建时间</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody emptyContent="暂无用户分组" items={userGroups}>
{(item) => (
<TableRow key={item.id}>
<TableCell>{item.name}</TableCell>
<TableCell>
{item.userNames.length > 0
? item.userNames.join("、")
: "-"}
</TableCell>
<TableCell>
<Chip
color={item.status === 1 ? "success" : "danger"}
size="sm"
>
{item.status === 1 ? "启用" : "停用"}
</Chip>
</TableCell>
<TableCell>{formatDate(item.createdTime)}</TableCell>
<TableCell>
<div className="flex gap-2">
<Button
size="sm"
variant="flat"
onPress={() => openAssignUsers(item)}
>
分配用户
</Button>
<Button
size="sm"
variant="light"
onPress={() => openEditUserGroup(item)}
>
编辑
</Button>
<Button
color="danger"
size="sm"
variant="light"
onPress={() => handleDeleteUserGroup(item.id)}
>
删除
</Button>
</div>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</CardBody>
</Card>
<Card>
<CardHeader>
<h3 className="text-lg font-semibold">权限分配</h3>
</CardHeader>
<CardBody className="space-y-4">
<div className="grid grid-cols-1 gap-3 md:grid-cols-3 md:items-end">
<Select
items={userGroups}
label="用户分组"
selectedKeys={
selectedUserGroupId ? [String(selectedUserGroupId)] : []
}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
setSelectedUserGroupId(key ? Number(key) : null);
}}
>
{(item) => <SelectItem key={item.id}>{item.name}</SelectItem>}
</Select>
<Select
items={tunnelGroups}
label="隧道分组"
selectedKeys={
selectedTunnelGroupId ? [String(selectedTunnelGroupId)] : []
}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
setSelectedTunnelGroupId(key ? Number(key) : null);
}}
>
{(item) => <SelectItem key={item.id}>{item.name}</SelectItem>}
</Select>
<Button
className="md:self-end md:justify-self-start whitespace-nowrap px-4"
color="primary"
isLoading={savingPermission}
size="sm"
onPress={handleAssignPermission}
>
分配
</Button>
</div>
<Table
aria-label="分组权限列表"
classNames={{
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
}}
>
<TableHeader>
<TableColumn>ID</TableColumn>
<TableColumn>用户分组</TableColumn>
<TableColumn>隧道分组</TableColumn>
<TableColumn>创建时间</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody emptyContent="暂无权限分配记录" items={permissions}>
{(item) => (
<TableRow key={item.id}>
<TableCell>{item.id}</TableCell>
<TableCell>
{item.userGroupName || item.userGroupId}
</TableCell>
<TableCell>
{item.tunnelGroupName || item.tunnelGroupId}
</TableCell>
<TableCell>{formatDate(item.createdTime)}</TableCell>
<TableCell>
<Button
color="danger"
size="sm"
variant="light"
onPress={() => handleRemovePermission(item.id)}
>
回收
</Button>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</CardBody>
</Card>
<Modal
backdrop="blur"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={tunnelGroupModalOpen}
onOpenChange={onTunnelGroupModalChange}
>
<TableHeader>
<TableColumn>名称</TableColumn>
<TableColumn>隧道</TableColumn>
<TableColumn>状态</TableColumn>
<TableColumn>创建时间</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody emptyContent="暂无隧道分组" items={tunnelGroups}>
{(item) => (
<TableRow key={item.id}>
<TableCell>{item.name}</TableCell>
<TableCell>
{item.tunnelNames.length > 0
? item.tunnelNames.join("、")
: "-"}
</TableCell>
<TableCell>
<Chip
color={item.status === 1 ? "success" : "danger"}
size="sm"
>
{item.status === 1 ? "启用" : "停用"}
</Chip>
</TableCell>
<TableCell>{formatDate(item.createdTime)}</TableCell>
<TableCell>
<div className="flex gap-2">
<Button
size="sm"
variant="flat"
onPress={() => openAssignTunnels(item)}
>
分配隧道
</Button>
<Button
size="sm"
variant="light"
onPress={() => openEditTunnelGroup(item)}
>
编辑
</Button>
<Button
color="danger"
size="sm"
variant="light"
onPress={() => handleDeleteTunnelGroup(item.id)}
>
删除
</Button>
</div>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</CardBody>
</Card>
<ModalContent>
<ModalHeader>
{editingTunnelGroup ? "编辑隧道分组" : "新建隧道分组"}
</ModalHeader>
<ModalBody className="space-y-3">
<Input
label="分组名称"
value={groupName}
onChange={(e) => setGroupName(e.target.value)}
/>
<Select
label="状态"
selectedKeys={[groupStatus]}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
<Card>
<CardHeader className="flex flex-row items-center gap-3 pb-2">
<h3 className="text-lg font-semibold">用户分组</h3>
<Button
className="h-7 px-3 text-xs font-medium min-w-0 shadow-sm"
color="primary"
size="sm"
onPress={openCreateUserGroup}
>
新建
</Button>
</CardHeader>
<CardBody>
<Table
aria-label="用户分组列表"
if (key) {
setGroupStatus(String(key));
}
}}
>
<SelectItem key="1">启用</SelectItem>
<SelectItem key="0">停用</SelectItem>
</Select>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onTunnelGroupModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingGroup}
onPress={saveTunnelGroup}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
<Modal
backdrop="blur"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={userGroupModalOpen}
onOpenChange={onUserGroupModalChange}
>
<TableHeader>
<TableColumn>名称</TableColumn>
<TableColumn>用户</TableColumn>
<TableColumn>状态</TableColumn>
<TableColumn>创建时间</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody emptyContent="暂无用户分组" items={userGroups}>
{(item) => (
<TableRow key={item.id}>
<TableCell>{item.name}</TableCell>
<TableCell>
{item.userNames.length > 0
? item.userNames.join("、")
: "-"}
</TableCell>
<TableCell>
<Chip
color={item.status === 1 ? "success" : "danger"}
size="sm"
>
{item.status === 1 ? "启用" : "停用"}
</Chip>
</TableCell>
<TableCell>{formatDate(item.createdTime)}</TableCell>
<TableCell>
<div className="flex gap-2">
<Button
size="sm"
variant="flat"
onPress={() => openAssignUsers(item)}
>
分配用户
</Button>
<Button
size="sm"
variant="light"
onPress={() => openEditUserGroup(item)}
>
编辑
</Button>
<Button
color="danger"
size="sm"
variant="light"
onPress={() => handleDeleteUserGroup(item.id)}
>
删除
</Button>
</div>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</CardBody>
</Card>
<ModalContent>
<ModalHeader>
{editingUserGroup ? "编辑用户分组" : "新建用户分组"}
</ModalHeader>
<ModalBody className="space-y-3">
<Input
label="分组名称"
value={groupName}
onChange={(e) => setGroupName(e.target.value)}
/>
<Select
label="状态"
selectedKeys={[groupStatus]}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
<Card>
<CardHeader>
<h3 className="text-lg font-semibold">权限分配</h3>
</CardHeader>
<CardBody className="space-y-4">
<div className="grid grid-cols-1 gap-3 md:grid-cols-3 md:items-end">
<Select
items={userGroups}
label="用户分组"
selectedKeys={
selectedUserGroupId ? [String(selectedUserGroupId)] : []
}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
if (key) {
setGroupStatus(String(key));
}
}}
>
<SelectItem key="1">启用</SelectItem>
<SelectItem key="0">停用</SelectItem>
</Select>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onUserGroupModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingGroup}
onPress={saveUserGroup}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
setSelectedUserGroupId(key ? Number(key) : null);
}}
>
{(item) => <SelectItem key={item.id}>{item.name}</SelectItem>}
</Select>
<Select
items={tunnelGroups}
label="隧道分组"
selectedKeys={
selectedTunnelGroupId ? [String(selectedTunnelGroupId)] : []
}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
setSelectedTunnelGroupId(key ? Number(key) : null);
}}
>
{(item) => <SelectItem key={item.id}>{item.name}</SelectItem>}
</Select>
<Button
className="md:self-end md:justify-self-start whitespace-nowrap px-4"
color="primary"
isLoading={savingPermission}
size="sm"
onPress={handleAssignPermission}
>
分配
</Button>
</div>
<Table
aria-label="分组权限列表"
<Modal
backdrop="blur"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={tunnelAssignModalOpen}
onOpenChange={onTunnelAssignModalChange}
>
<TableHeader>
<TableColumn>ID</TableColumn>
<TableColumn>用户分组</TableColumn>
<TableColumn>隧道分组</TableColumn>
<TableColumn>创建时间</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody emptyContent="暂无权限分配记录" items={permissions}>
{(item) => (
<TableRow key={item.id}>
<TableCell>{item.id}</TableCell>
<TableCell>
{item.userGroupName || item.userGroupId}
</TableCell>
<TableCell>
{item.tunnelGroupName || item.tunnelGroupId}
</TableCell>
<TableCell>{formatDate(item.createdTime)}</TableCell>
<TableCell>
<Button
color="danger"
size="sm"
variant="light"
onPress={() => handleRemovePermission(item.id)}
>
回收
</Button>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</CardBody>
</Card>
<ModalContent>
<ModalHeader>分配隧道 - {assignTunnelGroup?.name}</ModalHeader>
<ModalBody className="min-w-0">
<Select
className="min-w-0"
classNames={{ trigger: "max-w-full" }}
items={tunnels}
label="选择隧道"
selectedKeys={selectedTunnelKeys}
selectionMode="multiple"
onSelectionChange={(keys) => {
setSelectedTunnelKeys(
new Set(Array.from(keys as Set<React.Key>).map(String)),
);
}}
>
{(item) => <SelectItem key={item.id}>{item.name}</SelectItem>}
</Select>
<p
className="w-full min-w-0 max-w-full text-xs text-default-500 truncate"
title={`当前已选:${selectedTunnelSummary}`}
>
当前已选:{selectedTunnelSummary}
</p>
<p className="text-xs text-default-500">
不选择任何隧道并保存将清空该分组成员。
</p>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onTunnelAssignModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingAssign}
onPress={saveAssignTunnels}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
<Modal
backdrop="blur"
classNames={{
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={tunnelGroupModalOpen}
onOpenChange={onTunnelGroupModalChange}
>
<ModalContent>
<ModalHeader>
{editingTunnelGroup ? "编辑隧道分组" : "新建隧道分组"}
</ModalHeader>
<ModalBody className="space-y-3">
<Input
label="分组名称"
value={groupName}
onChange={(e) => setGroupName(e.target.value)}
/>
<Select
label="状态"
selectedKeys={[groupStatus]}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
if (key) {
setGroupStatus(String(key));
}
}}
>
<SelectItem key="1">启用</SelectItem>
<SelectItem key="0">停用</SelectItem>
</Select>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onTunnelGroupModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingGroup}
onPress={saveTunnelGroup}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
<Modal
backdrop="blur"
classNames={{
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={userGroupModalOpen}
onOpenChange={onUserGroupModalChange}
>
<ModalContent>
<ModalHeader>
{editingUserGroup ? "编辑用户分组" : "新建用户分组"}
</ModalHeader>
<ModalBody className="space-y-3">
<Input
label="分组名称"
value={groupName}
onChange={(e) => setGroupName(e.target.value)}
/>
<Select
label="状态"
selectedKeys={[groupStatus]}
onSelectionChange={(keys) => {
const key = Array.from(keys as Set<React.Key>)[0];
if (key) {
setGroupStatus(String(key));
}
}}
>
<SelectItem key="1">启用</SelectItem>
<SelectItem key="0">停用</SelectItem>
</Select>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onUserGroupModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingGroup}
onPress={saveUserGroup}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
<Modal
backdrop="blur"
classNames={{
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={tunnelAssignModalOpen}
onOpenChange={onTunnelAssignModalChange}
>
<ModalContent>
<ModalHeader>分配隧道 - {assignTunnelGroup?.name}</ModalHeader>
<ModalBody className="min-w-0">
<Select
className="min-w-0"
classNames={{ trigger: "max-w-full" }}
items={tunnels}
label="选择隧道"
selectedKeys={selectedTunnelKeys}
selectionMode="multiple"
onSelectionChange={(keys) => {
setSelectedTunnelKeys(
new Set(Array.from(keys as Set<React.Key>).map(String)),
);
}}
>
{(item) => <SelectItem key={item.id}>{item.name}</SelectItem>}
</Select>
<p
className="w-full min-w-0 max-w-full text-xs text-default-500 truncate"
title={`当前已选:${selectedTunnelSummary}`}
>
当前已选:{selectedTunnelSummary}
</p>
<p className="text-xs text-default-500">
不选择任何隧道并保存将清空该分组成员。
</p>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onTunnelAssignModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingAssign}
onPress={saveAssignTunnels}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
<Modal
backdrop="blur"
classNames={{
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={userAssignModalOpen}
onOpenChange={onUserAssignModalChange}
>
<ModalContent>
<ModalHeader>分配用户 - {assignUserGroup?.name}</ModalHeader>
<ModalBody className="min-w-0">
<Select
className="min-w-0"
classNames={{ trigger: "max-w-full" }}
items={users}
label="选择用户"
selectedKeys={selectedUserKeys}
selectionMode="multiple"
onSelectionChange={(keys) => {
setSelectedUserKeys(
new Set(Array.from(keys as Set<React.Key>).map(String)),
);
}}
>
{(item) => <SelectItem key={item.id}>{item.user}</SelectItem>}
</Select>
<p
className="w-full min-w-0 max-w-full text-xs text-default-500 truncate"
title={`当前已选:${selectedUserSummary}`}
>
当前已选:{selectedUserSummary}
</p>
<p className="text-xs text-default-500">
不选择任何用户并保存将清空该分组成员。
</p>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onUserAssignModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingAssign}
onPress={saveAssignUsers}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
<Modal
backdrop="blur"
classNames={{
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
}}
isOpen={userAssignModalOpen}
onOpenChange={onUserAssignModalChange}
>
<ModalContent>
<ModalHeader>分配用户 - {assignUserGroup?.name}</ModalHeader>
<ModalBody className="min-w-0">
<Select
className="min-w-0"
classNames={{ trigger: "max-w-full" }}
items={users}
label="选择用户"
selectedKeys={selectedUserKeys}
selectionMode="multiple"
onSelectionChange={(keys) => {
setSelectedUserKeys(
new Set(Array.from(keys as Set<React.Key>).map(String)),
);
}}
>
{(item) => <SelectItem key={item.id}>{item.user}</SelectItem>}
</Select>
<p
className="w-full min-w-0 max-w-full text-xs text-default-500 truncate"
title={`当前已选:${selectedUserSummary}`}
>
当前已选:{selectedUserSummary}
</p>
<p className="text-xs text-default-500">
不选择任何用户并保存将清空该分组成员。
</p>
</ModalBody>
<ModalFooter>
<Button variant="light" onPress={onUserAssignModalClose}>
取消
</Button>
<Button
color="primary"
isLoading={savingAssign}
onPress={saveAssignUsers}
>
保存
</Button>
</ModalFooter>
</ModalContent>
</Modal>
</>
)}
</AnimatedPage>
+12 -4
View File
@@ -162,7 +162,9 @@ export default function IndexPage() {
<Card className="w-full bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_20px_40px_rgba(0,0,0,0.15)] border-white/80 dark:border-white/10 rounded-[32px] p-2 sm:p-4">
<CardHeader className="pb-0 pt-6 px-6 flex-col items-center">
<BrandLogo className="w-14 h-14 rounded-2xl mb-4" size={56} />
<h1 className="text-2xl font-bold tracking-tight text-foreground">{siteConfig.name}</h1>
<h1 className="text-2xl font-bold tracking-tight text-foreground">
{siteConfig.name}
</h1>
<p className="text-sm text-default-500 mt-2 font-medium">
Sign in to manage your networks
</p>
@@ -171,7 +173,8 @@ export default function IndexPage() {
<div className="flex flex-col gap-5">
<Input
classNames={{
inputWrapper: "bg-white/50 dark:bg-black/50 backdrop-blur-md border border-white/60 dark:border-white/10 h-12 shadow-sm rounded-xl",
inputWrapper:
"bg-white/10 dark:bg-white/5 backdrop-blur-md border border-white/30 dark:border-white/10 h-12 shadow-sm rounded-xl",
input: "text-base font-medium",
}}
errorMessage={errors.username}
@@ -189,7 +192,8 @@ export default function IndexPage() {
<Input
classNames={{
inputWrapper: "bg-white/50 dark:bg-black/50 backdrop-blur-md border border-white/60 dark:border-white/10 h-12 shadow-sm rounded-xl",
inputWrapper:
"bg-white/10 dark:bg-white/5 backdrop-blur-md border border-white/30 dark:border-white/10 h-12 shadow-sm rounded-xl",
input: "text-base font-medium tracking-wider",
}}
isDisabled={loading}
@@ -211,7 +215,11 @@ export default function IndexPage() {
isLoading={loading}
onPress={handleLogin}
>
{loading ? (showCaptcha ? "Verifying..." : "Signing in...") : "Sign In"}
{loading
? showCaptcha
? "Verifying..."
: "Signing in..."
: "Sign In"}
</Button>
</div>
</CardBody>
+140 -135
View File
@@ -1,5 +1,6 @@
import { useState, useEffect, useMemo } from "react";
import toast from "react-hot-toast";
import { LayoutGrid, List } from "lucide-react";
import {
AnimatedPage,
@@ -18,7 +19,6 @@ import {
ModalFooter,
} from "@/shadcn-bridge/heroui/modal";
import { Chip } from "@/shadcn-bridge/heroui/chip";
import { LayoutGrid, List } from "lucide-react";
import {
Table,
TableHeader,
@@ -251,7 +251,11 @@ export default function LimitPage() {
variant="flat"
onPress={() => setViewMode(viewMode === "list" ? "grid" : "list")}
>
{viewMode === "list" ? <LayoutGrid className="w-4 h-4" /> : <List className="w-4 h-4" />}
{viewMode === "list" ? (
<LayoutGrid className="w-4 h-4" />
) : (
<List className="w-4 h-4" />
)}
</Button>
<Button color="primary" size="sm" variant="flat" onPress={handleAdd}>
新增
@@ -263,153 +267,154 @@ export default function LimitPage() {
{filteredRules.length > 0 ? (
viewMode === "list" ? (
<Card>
<Table
aria-label="限速规则列表"
className="overflow-x-auto min-w-full"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-[24px]",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
}}
>
<TableHeader>
<TableColumn>规则名称</TableColumn>
<TableColumn>速度限制</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody items={filteredRules}>
{(rule) => (
<TableRow key={rule.id}>
<TableCell>
<div className="flex items-center gap-2">
<div
className={`shrink-0 w-2 h-2 rounded-full ${
rule.status === 1 ? "bg-success" : "bg-danger"
}`}
/>
<span className="font-medium text-foreground text-sm">
{rule.name}
</span>
</div>
</TableCell>
<TableCell>
<span className="text-sm font-mono text-default-600">
{rule.speed} Mbps
</span>
</TableCell>
<TableCell>
<div className="flex flex-wrap items-center gap-1.5 min-w-max">
<Button
className="h-6 px-2 min-w-0 text-xs bg-indigo-50 text-indigo-600 hover:bg-indigo-100 dark:bg-indigo-950/30 dark:text-indigo-400"
<Table
aria-label="限速规则列表"
className="overflow-x-auto min-w-full"
classNames={{
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-[24px]",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
}}
>
<TableHeader>
<TableColumn>规则名称</TableColumn>
<TableColumn>速度限制</TableColumn>
<TableColumn>操作</TableColumn>
</TableHeader>
<TableBody items={filteredRules}>
{(rule) => (
<TableRow key={rule.id}>
<TableCell>
<div className="flex items-center gap-2">
<div
className={`shrink-0 w-2 h-2 rounded-full ${
rule.status === 1 ? "bg-success" : "bg-danger"
}`}
/>
<span className="font-medium text-foreground text-sm">
{rule.name}
</span>
</div>
</TableCell>
<TableCell>
<span className="text-sm font-mono text-default-600">
{rule.speed} Mbps
</span>
</TableCell>
<TableCell>
<div className="flex flex-wrap items-center gap-1.5 min-w-max">
<Button
className="h-6 px-2 min-w-0 text-xs bg-indigo-50 text-indigo-600 hover:bg-indigo-100 dark:bg-indigo-950/30 dark:text-indigo-400"
size="sm"
variant="flat"
onPress={() => handleEdit(rule)}
>
编辑
</Button>
<Button
className="h-6 px-2 min-w-0 text-xs bg-rose-50 text-rose-600 hover:bg-rose-100 dark:bg-rose-950/30 dark:text-rose-400"
size="sm"
variant="flat"
onPress={() => handleDelete(rule)}
>
删除
</Button>
</div>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</Card>
) : (
<StaggerList className="grid grid-cols-1 sm:grid-cols-2 lg:grid-cols-3 xl:grid-cols-4 2xl:grid-cols-5 gap-4">
{filteredRules.map((rule) => (
<StaggerItem key={rule.id}>
<Card className="shadow-sm overflow-hidden h-full">
<CardHeader className="pb-2 md:pb-2">
<div className="flex justify-between items-start w-full">
<div>
<h3 className="font-semibold text-foreground">
{rule.name}
</h3>
</div>
<Chip
color={rule.status === 1 ? "success" : "danger"}
size="sm"
variant="flat"
>
{rule.status === 1 ? "运行" : "异常"}
</Chip>
</div>
</CardHeader>
<CardBody className="pt-0 pb-3 md:pt-0 md:pb-3">
<div className="space-y-3">
<div className="flex justify-between items-center">
<span className="text-small text-default-600">
速度限制
</span>
<Chip color="secondary" size="sm" variant="flat">
{rule.speed} Mbps
</Chip>
</div>
</div>
<div className="flex gap-2 mt-4">
<Button
className="flex-1"
color="primary"
size="sm"
startContent={
<svg
aria-hidden="true"
className="w-4 h-4"
fill="currentColor"
viewBox="0 0 20 20"
>
<path d="M13.586 3.586a2 2 0 112.828 2.828l-.793.793-2.828-2.828.793-.793zM11.379 5.793L3 14.172V17h2.828l8.38-8.379-2.83-2.828z" />
</svg>
}
variant="flat"
onPress={() => handleEdit(rule)}
>
编辑
</Button>
<Button
className="h-6 px-2 min-w-0 text-xs bg-rose-50 text-rose-600 hover:bg-rose-100 dark:bg-rose-950/30 dark:text-rose-400"
className="flex-1"
color="danger"
size="sm"
startContent={
<svg
aria-hidden="true"
className="w-4 h-4"
fill="currentColor"
viewBox="0 0 20 20"
>
<path
clipRule="evenodd"
d="M9 2a1 1 0 000 2h2a1 1 0 100-2H9z"
fillRule="evenodd"
/>
<path
clipRule="evenodd"
d="M10 18a8 8 0 100-16 8 8 0 000 16zM8 7a1 1 0 012 0v4a1 1 0 11-2 0V7zM12 7a1 1 0 012 0v4a1 1 0 11-2 0V7z"
fillRule="evenodd"
/>
</svg>
}
variant="flat"
onPress={() => handleDelete(rule)}
>
删除
</Button>
</div>
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</Card>
) : (
<StaggerList className="grid grid-cols-1 sm:grid-cols-2 lg:grid-cols-3 xl:grid-cols-4 2xl:grid-cols-5 gap-4">
{filteredRules.map((rule) => (
<StaggerItem key={rule.id}>
<Card className="shadow-sm overflow-hidden h-full">
<CardHeader className="pb-2 md:pb-2">
<div className="flex justify-between items-start w-full">
<div>
<h3 className="font-semibold text-foreground">
{rule.name}
</h3>
</div>
<Chip
color={rule.status === 1 ? "success" : "danger"}
size="sm"
variant="flat"
>
{rule.status === 1 ? "运行" : "异常"}
</Chip>
</div>
</CardHeader>
<CardBody className="pt-0 pb-3 md:pt-0 md:pb-3">
<div className="space-y-3">
<div className="flex justify-between items-center">
<span className="text-small text-default-600">
速度限制
</span>
<Chip color="secondary" size="sm" variant="flat">
{rule.speed} Mbps
</Chip>
</div>
</div>
<div className="flex gap-2 mt-4">
<Button
className="flex-1"
color="primary"
size="sm"
startContent={
<svg
aria-hidden="true"
className="w-4 h-4"
fill="currentColor"
viewBox="0 0 20 20"
>
<path d="M13.586 3.586a2 2 0 112.828 2.828l-.793.793-2.828-2.828.793-.793zM11.379 5.793L3 14.172V17h2.828l8.38-8.379-2.83-2.828z" />
</svg>
}
variant="flat"
onPress={() => handleEdit(rule)}
>
编辑
</Button>
<Button
className="flex-1"
color="danger"
size="sm"
startContent={
<svg
aria-hidden="true"
className="w-4 h-4"
fill="currentColor"
viewBox="0 0 20 20"
>
<path
clipRule="evenodd"
d="M9 2a1 1 0 000 2h2a1 1 0 100-2H9z"
fillRule="evenodd"
/>
<path
clipRule="evenodd"
d="M10 18a8 8 0 100-16 8 8 0 000 16zM8 7a1 1 0 012 0v4a1 1 0 11-2 0V7zM12 7a1 1 0 012 0v4a1 1 0 11-2 0V7z"
fillRule="evenodd"
/>
</svg>
}
variant="flat"
onPress={() => handleDelete(rule)}
>
删除
</Button>
</div>
</CardBody>
</Card>
</StaggerItem>
))}
</StaggerList>
</CardBody>
</Card>
</StaggerItem>
))}
</StaggerList>
)
) : (
/* 空状态 */
+47 -12
View File
@@ -2,7 +2,13 @@ import type { MonitorNodeApiItem } from "@/api/types";
import { useCallback, useEffect, useMemo, useState } from "react";
import toast from "react-hot-toast";
import { RefreshCw, LayoutGrid, List, Server, ArrowRightLeft } from "lucide-react";
import {
RefreshCw,
LayoutGrid,
List,
Server,
ArrowRightLeft,
} from "lucide-react";
import { AnimatedPage } from "@/components/animated-page";
import { Button } from "@/shadcn-bridge/heroui/button";
@@ -38,11 +44,16 @@ export default function MonitorPage() {
const [viewMode, setViewMode] = useState<"list" | "grid">("list");
const [activeTab, setActiveTab] = useState<MonitorTab>("nodes");
const [realtimeNodeStatus, setRealtimeNodeStatus] = useState<Record<number, "online" | "offline">>({});
const [realtimeNodeMetrics, setRealtimeNodeMetrics] = useState<Record<number, any>>({});
const [realtimeNodeStatus, setRealtimeNodeStatus] = useState<
Record<number, "online" | "offline">
>({});
const [realtimeNodeMetrics, setRealtimeNodeMetrics] = useState<
Record<number, any>
>({});
const handleRealtimeMessage = useCallback((message: any) => {
const nodeId = Number(message?.id ?? 0);
if (!nodeId || Number.isNaN(nodeId)) return;
const type = String(message?.type ?? "");
@@ -50,15 +61,18 @@ export default function MonitorPage() {
if (type === "status") {
const status = Number(payload);
setRealtimeNodeStatus((prev) => ({
...prev,
[nodeId]: status === 1 ? "online" : "offline",
}));
return;
}
if (type === "metric") {
let raw = payload;
if (typeof raw === "string") {
try {
raw = JSON.parse(raw);
@@ -77,7 +91,7 @@ export default function MonitorPage() {
netOutSpeed: Number(metric.netOutSpeed ?? metric.net_out_speed ?? 0),
tcpConns: Number(metric.tcpConns ?? metric.tcp_conns ?? 0),
udpConns: Number(metric.udpConns ?? metric.udp_conns ?? 0),
}
},
}));
}
}, []);
@@ -89,6 +103,7 @@ export default function MonitorPage() {
const loadNodes = useCallback(async (options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
if (!silent) setNodesLoading(true);
try {
const response = await getMonitorNodes();
@@ -133,7 +148,9 @@ export default function MonitorPage() {
.map((n) => ({
id: Number(n.id),
name: String(n.name ?? ""),
connectionStatus: realtimeNodeStatus[Number(n.id)] || (n.status === 1 ? "online" : "offline"),
connectionStatus:
realtimeNodeStatus[Number(n.id)] ||
(n.status === 1 ? "online" : "offline"),
version: n.version,
}));
@@ -151,6 +168,7 @@ export default function MonitorPage() {
if (node.connectionStatus === "online") {
onlineCount++;
const metric = realtimeNodeMetrics[node.id];
if (metric) {
totalCpu += metric.cpuUsage;
totalConns += metric.tcpConns + metric.udpConns;
@@ -161,6 +179,7 @@ export default function MonitorPage() {
});
const avgCpu = onlineCount > 0 ? totalCpu / onlineCount : 0;
return {
avgCpu,
totalConns,
@@ -174,10 +193,14 @@ export default function MonitorPage() {
<div className="grid grid-cols-1 md:grid-cols-3 gap-4 mb-8">
<div className="rounded-3xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_10px_30px_rgba(0,0,0,0.1)] p-6 relative overflow-hidden flex flex-col justify-between h-40">
<div className="flex justify-between items-center z-10 relative">
<span className="text-default-600 font-medium text-sm">System Load</span>
<span className="text-default-600 font-medium text-sm">
System Load
</span>
</div>
<div className="z-10 relative">
<span className="text-4xl font-bold text-foreground">{aggregateMetrics.avgCpu.toFixed(1)}%</span>
<span className="text-4xl font-bold text-foreground">
{aggregateMetrics.avgCpu.toFixed(1)}%
</span>
</div>
<div className="absolute bottom-0 left-0 right-0 h-12 flex items-end gap-1 px-6 pb-4 opacity-50 z-0">
<div className="w-full bg-primary/40 h-2 rounded-t-sm" />
@@ -191,10 +214,14 @@ export default function MonitorPage() {
<div className="rounded-3xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_10px_30px_rgba(0,0,0,0.1)] p-6 relative overflow-hidden flex flex-col justify-between h-40">
<div className="flex justify-between items-center z-10 relative">
<span className="text-default-600 font-medium text-sm">Active Connections</span>
<span className="text-default-600 font-medium text-sm">
Active Connections
</span>
</div>
<div className="z-10 relative">
<span className="text-4xl font-bold text-foreground">{aggregateMetrics.totalConns}</span>
<span className="text-4xl font-bold text-foreground">
{aggregateMetrics.totalConns}
</span>
</div>
<div className="absolute bottom-0 left-0 right-0 h-12 flex items-end gap-1 px-6 pb-4 opacity-50 z-0">
<div className="w-full bg-success/40 h-3.5 rounded-t-sm" />
@@ -208,10 +235,14 @@ export default function MonitorPage() {
<div className="rounded-3xl border border-white/80 dark:border-white/10 bg-white/20 dark:bg-zinc-900/20 backdrop-blur-3xl shadow-[0_10px_30px_rgba(0,0,0,0.1)] p-6 relative overflow-hidden flex flex-col justify-between h-40">
<div className="flex justify-between items-center z-10 relative">
<span className="text-default-600 font-medium text-sm">Bandwidth</span>
<span className="text-default-600 font-medium text-sm">
Bandwidth
</span>
</div>
<div className="z-10 relative">
<span className="text-4xl font-bold text-foreground">{formatBytesPerSecond(aggregateMetrics.totalBandwidth)}</span>
<span className="text-4xl font-bold text-foreground">
{formatBytesPerSecond(aggregateMetrics.totalBandwidth)}
</span>
</div>
<div className="absolute bottom-0 left-0 right-0 h-12 flex items-end gap-1 px-6 pb-4 opacity-50 z-0">
<div className="w-full bg-secondary/40 h-2 rounded-t-sm" />
@@ -236,7 +267,11 @@ export default function MonitorPage() {
variant="flat"
onPress={() => setViewMode(viewMode === "list" ? "grid" : "list")}
>
{viewMode === "list" ? <LayoutGrid className="w-4 h-4" /> : <List className="w-4 h-4" />}
{viewMode === "list" ? (
<LayoutGrid className="w-4 h-4" />
) : (
<List className="w-4 h-4" />
)}
</Button>
{activeTab === "nodes" && (
<Button
+36 -21
View File
@@ -17,10 +17,10 @@ import {
useSortable,
} from "@dnd-kit/sortable";
import { CSS } from "@dnd-kit/utilities";
import { LayoutGrid, List } from "lucide-react";
import { SearchBar } from "@/components/search-bar";
import { AnimatedPage } from "@/components/animated-page";
import { LayoutGrid, List } from "lucide-react";
import {
Table,
TableHeader,
@@ -699,7 +699,6 @@ export default function NodePage() {
// 格式化流量
const formatFlow = (bytes: number): string => {
if (!Number.isFinite(bytes) || bytes <= 0) {
return "0 B";
@@ -726,7 +725,6 @@ export default function NodePage() {
return "未知链路";
};
// IPv4/IPv6 格式验证(仅用于判定地址族)
const ipv4Regex =
/^(25[0-5]|2[0-4][0-9]|[01]?[0-9][0-9]?)\.(25[0-5]|2[0-4][0-9]|[01]?[0-9][0-9]?)\.(25[0-5]|2[0-4][0-9]|[01]?[0-9][0-9]?)\.(25[0-5]|2[0-4][0-9]|[01]?[0-9][0-9]?)$/;
@@ -1650,10 +1648,12 @@ export default function NodePage() {
{/* 视图切换按钮 */}
<Button
isIconOnly
className="text-default-600 hidden sm:flex"
size="sm"
variant="flat"
className="text-default-600 hidden sm:flex"
onPress={() => setViewMode(viewMode === "list" ? "grid" : "list")}
onPress={() =>
setViewMode(viewMode === "list" ? "grid" : "list")
}
>
{viewMode === "list" ? (
<LayoutGrid className="w-4 h-4" />
@@ -1767,7 +1767,8 @@ export default function NodePage() {
aria-label="节点列表"
className="overflow-x-auto min-w-full"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-2xl",
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/10 dark:hover:bg-white/5 transition-colors",
@@ -1776,7 +1777,11 @@ export default function NodePage() {
<TableHeader>
<TableColumn className="w-12 px-4 whitespace-nowrap overflow-hidden">
<Checkbox
isSelected={selectMode && selectedIds.size === displayNodes.length && displayNodes.length > 0}
isSelected={
selectMode &&
selectedIds.size === displayNodes.length &&
displayNodes.length > 0
}
onValueChange={(checked) => {
if (checked) {
selectAll();
@@ -1806,12 +1811,16 @@ export default function NodePage() {
onValueChange={(checked) => {
if (checked) {
setSelectMode(true);
setSelectedIds((prev) => new Set([...prev, node.id]));
setSelectedIds(
(prev) => new Set([...prev, node.id]),
);
} else {
setSelectedIds((prev) => {
const next = new Set(prev);
next.delete(node.id);
if (next.size === 0) setSelectMode(false);
return next;
});
}
@@ -1833,7 +1842,11 @@ export default function NodePage() {
{node.name}
</span>
{hasRemark && (
<Chip size="sm" variant="flat" className="h-4 px-1 text-[10px]">
<Chip
className="h-4 px-1 text-[10px]"
size="sm"
variant="flat"
>
{node.remark}
</Chip>
)}
@@ -1841,7 +1854,9 @@ export default function NodePage() {
</TableCell>
<TableCell>
<span className="text-sm font-mono text-default-600">
{isRemoteNode ? new URL(node.remoteUrl || "").hostname : node.serverIp}
{isRemoteNode
? new URL(node.remoteUrl || "").hostname
: node.serverIp}
</span>
</TableCell>
<TableCell>
@@ -1853,10 +1868,10 @@ export default function NodePage() {
<div className="flex flex-wrap items-center gap-1.5 min-w-max">
{!isRemoteNode && (
<Button
size="sm"
variant="flat"
className="h-6 px-2 min-w-0 text-xs bg-emerald-50 text-emerald-600 hover:bg-emerald-100 dark:bg-emerald-950/30 dark:text-emerald-400"
isLoading={node.copyLoading}
size="sm"
variant="flat"
onPress={() => openInstallSelector(node)}
>
安装
@@ -1864,11 +1879,11 @@ export default function NodePage() {
)}
{!isRemoteNode && (
<Button
size="sm"
variant="flat"
className="h-6 px-2 min-w-0 text-xs bg-amber-50 text-amber-600 hover:bg-amber-100 dark:bg-amber-950/30 dark:text-amber-400"
isDisabled={node.connectionStatus !== "online"}
isLoading={node.upgradeLoading}
size="sm"
variant="flat"
onPress={() => {
setUpgradeTargetNodeId(node.id);
openUpgradeModal("single");
@@ -1879,11 +1894,11 @@ export default function NodePage() {
)}
{!isRemoteNode && (
<Button
size="sm"
variant="flat"
className="h-6 px-2 min-w-0 text-xs bg-blue-50 text-blue-600 hover:bg-blue-100 dark:bg-blue-950/30 dark:text-blue-400"
isDisabled={node.connectionStatus !== "online"}
isLoading={node.rollbackLoading}
size="sm"
variant="flat"
onPress={() => handleRollbackNode(node)}
>
回退
@@ -1891,18 +1906,18 @@ export default function NodePage() {
)}
{!isRemoteNode && (
<Button
className="h-6 px-2 min-w-0 text-xs bg-indigo-50 text-indigo-600 hover:bg-indigo-100 dark:bg-indigo-950/30 dark:text-indigo-400"
size="sm"
variant="flat"
className="h-6 px-2 min-w-0 text-xs bg-indigo-50 text-indigo-600 hover:bg-indigo-100 dark:bg-indigo-950/30 dark:text-indigo-400"
onPress={() => handleEdit(node)}
>
编辑
</Button>
)}
<Button
className="h-6 px-2 min-w-0 text-xs bg-rose-50 text-rose-600 hover:bg-rose-100 dark:bg-rose-950/30 dark:text-rose-400"
size="sm"
variant="flat"
className="h-6 px-2 min-w-0 text-xs bg-rose-50 text-rose-600 hover:bg-rose-100 dark:bg-rose-950/30 dark:text-rose-400"
onPress={() => handleDelete(node)}
>
删除
@@ -2054,7 +2069,9 @@ export default function NodePage() {
type="button"
onClick={(e) => {
e.stopPropagation();
handleDismissExpiryReminder(node.id);
handleDismissExpiryReminder(
node.id,
);
}}
>
关闭提醒
@@ -2327,8 +2344,6 @@ export default function NodePage() {
</div>
)}
<div className="mt-auto space-y-3">
{/* 操作按钮 */}
<div className="space-y-1.5">
File diff suppressed because it is too large Load Diff
@@ -5,7 +5,13 @@ import type {
TunnelQualityHopApiItem,
} from "@/api/types";
import React, { useCallback, useEffect, useMemo, useRef, useState } from "react";
import React, {
useCallback,
useEffect,
useMemo,
useRef,
useState,
} from "react";
import {
LineChart,
Line,
@@ -35,7 +41,6 @@ import {
getMonitorTunnelQualityHistory,
getConfigByName,
} from "@/api";
import { Button } from "@/shadcn-bridge/heroui/button";
import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card";
import { Chip } from "@/shadcn-bridge/heroui/chip";
@@ -54,7 +59,8 @@ interface TunnelMonitorViewProps {
}
const QUALITY_POLL_INTERVAL = 1_000; // 1 second
const MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY = "monitor_tunnel_quality_enabled";
const MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY =
"monitor_tunnel_quality_enabled";
const MONITOR_TUNNEL_QUALITY_ENABLED_EVENT =
"monitorTunnelQualityEnabledChanged";
@@ -78,7 +84,13 @@ const formatTimestamp = (ts: number, rangeMs?: number): string => {
});
};
function LatencyDisplay({ value, loading }: { value?: number; loading?: boolean }) {
function LatencyDisplay({
value,
loading,
}: {
value?: number;
loading?: boolean;
}) {
if (loading) {
return <RefreshCw className="w-3 h-3 animate-spin inline text-primary" />;
}
@@ -87,11 +99,16 @@ function LatencyDisplay({ value, loading }: { value?: number; loading?: boolean
}
const ms = value.toFixed(0);
let colorClass = "text-success";
if (value > 200) colorClass = "text-danger";
else if (value > 100) colorClass = "text-warning";
else if (value > 50) colorClass = "text-primary";
return <span className={`font-mono text-xs font-semibold ${colorClass}`}>{ms}ms</span>;
return (
<span className={`font-mono text-xs font-semibold ${colorClass}`}>
{ms}ms
</span>
);
}
function LiveDot() {
@@ -116,9 +133,11 @@ function UptimeHistoryBar({
}) {
const bars = [...Array(UPTIME_KUMA_BARS_COUNT)].map((_, i) => {
const historyIndex = history.length - UPTIME_KUMA_BARS_COUNT + i;
if (historyIndex >= 0 && historyIndex < history.length) {
return history[historyIndex];
}
return null;
});
@@ -183,13 +202,20 @@ const TIME_RANGE_OPTIONS = [
{ key: String(24 * 60 * 60 * 1000), label: "24小时" },
];
function TimeRangeSelect({ value, onChange }: { value: number; onChange: (v: number) => void }) {
function TimeRangeSelect({
value,
onChange,
}: {
value: number;
onChange: (v: number) => void;
}) {
return (
<Select
className="w-36"
selectedKeys={[String(value)]}
onSelectionChange={(keys) => {
const v = Number(Array.from(keys)[0]);
if (v > 0) onChange(v);
}}
>
@@ -205,13 +231,23 @@ interface QualityChartCardProps {
onRangeChange: (v: number) => void;
loading: boolean;
error: string | null;
data: Array<{ time: string; entryToExit: number | null; exitToBing: number | null }>;
data: Array<{
time: string;
entryToExit: number | null;
exitToBing: number | null;
}>;
tunnelId: number;
onRefresh: (id: number) => void;
}
const QualityChartCard = React.memo(function QualityChartCard({
rangeMs, onRangeChange, loading, error, data, tunnelId, onRefresh,
rangeMs,
onRangeChange,
loading,
error,
data,
tunnelId,
onRefresh,
}: QualityChartCardProps) {
return (
<Card>
@@ -219,7 +255,12 @@ const QualityChartCard = React.memo(function QualityChartCard({
<h3 className="text-lg font-semibold">质量趋势</h3>
<div className="flex items-center gap-2">
<TimeRangeSelect value={rangeMs} onChange={onRangeChange} />
<Button isLoading={loading} size="sm" variant="flat" onPress={() => onRefresh(tunnelId)}>
<Button
isLoading={loading}
size="sm"
variant="flat"
onPress={() => onRefresh(tunnelId)}
>
<RefreshCw className="w-4 h-4 mr-1" />
刷新
</Button>
@@ -227,38 +268,75 @@ const QualityChartCard = React.memo(function QualityChartCard({
</CardHeader>
<CardBody>
{loading ? (
<div className="flex justify-center py-8"><RefreshCw className="w-6 h-6 animate-spin" /></div>
<div className="flex justify-center py-8">
<RefreshCw className="w-6 h-6 animate-spin" />
</div>
) : error ? (
<div className="text-center py-8 text-danger text-sm">{error}</div>
) : data.length > 0 ? (
<div className="h-64">
<ResponsiveContainer height="100%" width="100%">
<LineChart data={data}>
<CartesianGrid strokeDasharray="3 3" opacity={0.3} />
<CartesianGrid opacity={0.3} strokeDasharray="3 3" />
<XAxis dataKey="time" fontSize={11} tick={{ fill: "#888" }} />
<YAxis
fontSize={11}
label={{
value: "延迟 (ms)",
angle: -90,
position: "insideLeft",
style: { fontSize: 11, fill: "#888" },
}}
tick={{ fill: "#888" }}
tickFormatter={(v: any) => `${Number(v).toFixed(0)}ms`}
label={{ value: "延迟 (ms)", angle: -90, position: "insideLeft", style: { fontSize: 11, fill: "#888" } }}
/>
<Tooltip
contentStyle={{ backgroundColor: "rgba(0,0,0,0.85)", border: "none", borderRadius: "8px", fontSize: 12 }}
labelStyle={{ color: "#fff" }}
contentStyle={{
backgroundColor: "rgba(0,0,0,0.85)",
border: "none",
borderRadius: "8px",
fontSize: 12,
}}
formatter={(value: unknown, name: string) => {
const n = Number(value);
if (!Number.isFinite(n)) return "-";
const label = name === "entryToExit" ? "入口→出口" : name === "exitToBing" ? "出口→Bing" : name;
const label =
name === "entryToExit"
? "入口→出口"
: name === "exitToBing"
? "出口→Bing"
: name;
return [`${n.toFixed(1)}ms`, label];
}}
labelStyle={{ color: "#fff" }}
/>
<Line
connectNulls
dataKey="entryToExit"
dot={false}
name="entryToExit"
stroke="#10b981"
strokeWidth={2}
type="monotone"
/>
<Line
connectNulls
dataKey="exitToBing"
dot={false}
name="exitToBing"
stroke="#3b82f6"
strokeWidth={2}
type="monotone"
/>
<Line connectNulls dataKey="entryToExit" dot={false} name="entryToExit" stroke="#10b981" strokeWidth={2} type="monotone" />
<Line connectNulls dataKey="exitToBing" dot={false} name="exitToBing" stroke="#3b82f6" strokeWidth={2} type="monotone" />
</LineChart>
</ResponsiveContainer>
</div>
) : (
<div className="text-center py-8 text-default-500">暂无质量历史数据</div>
<div className="text-center py-8 text-default-500">
暂无质量历史数据
</div>
)}
</CardBody>
</Card>
@@ -270,20 +348,33 @@ interface TrafficChartCardProps {
onRangeChange: (v: number) => void;
loading: boolean;
error: string | null;
data: Array<{ time: string; bytesIn: number; bytesOut: number; connections: number }>;
data: Array<{
time: string;
bytesIn: number;
bytesOut: number;
connections: number;
}>;
tunnelId: number;
onRefresh: (id: number) => void;
}
const TrafficChartCard = React.memo(function TrafficChartCard({
rangeMs, onRangeChange, loading, error, data, tunnelId, onRefresh,
rangeMs,
onRangeChange,
loading,
error,
data,
tunnelId,
onRefresh,
}: TrafficChartCardProps) {
const yFormatter = (value: unknown) => {
const n = Number(value);
if (!Number.isFinite(n) || n <= 0) return "0 B";
const k = 1024;
const sizes = ["B", "KB", "MB", "GB", "TB"];
const i = Math.floor(Math.log(n) / Math.log(k));
return `${parseFloat((n / Math.pow(k, i)).toFixed(2))} ${sizes[i]}`;
};
@@ -293,7 +384,12 @@ const TrafficChartCard = React.memo(function TrafficChartCard({
<h3 className="text-lg font-semibold">流量趋势</h3>
<div className="flex items-center gap-2">
<TimeRangeSelect value={rangeMs} onChange={onRangeChange} />
<Button isLoading={loading} size="sm" variant="flat" onPress={() => onRefresh(tunnelId)}>
<Button
isLoading={loading}
size="sm"
variant="flat"
onPress={() => onRefresh(tunnelId)}
>
<RefreshCw className="w-4 h-4 mr-1" />
刷新
</Button>
@@ -301,23 +397,48 @@ const TrafficChartCard = React.memo(function TrafficChartCard({
</CardHeader>
<CardBody>
{loading ? (
<div className="flex justify-center py-8"><RefreshCw className="w-6 h-6 animate-spin" /></div>
<div className="flex justify-center py-8">
<RefreshCw className="w-6 h-6 animate-spin" />
</div>
) : error ? (
<div className="text-center py-8 text-danger text-sm">{error}</div>
) : data.length > 0 ? (
<div className="h-64">
<ResponsiveContainer height="100%" width="100%">
<LineChart data={data}>
<CartesianGrid strokeDasharray="3 3" opacity={0.3} />
<CartesianGrid opacity={0.3} strokeDasharray="3 3" />
<XAxis dataKey="time" fontSize={11} tick={{ fill: "#888" }} />
<YAxis fontSize={11} tick={{ fill: "#888" }} tickFormatter={yFormatter} />
<Tooltip
contentStyle={{ backgroundColor: "rgba(0,0,0,0.85)", border: "none", borderRadius: "8px", fontSize: 12 }}
labelStyle={{ color: "#fff" }}
formatter={yFormatter}
<YAxis
fontSize={11}
tick={{ fill: "#888" }}
tickFormatter={yFormatter}
/>
<Tooltip
contentStyle={{
backgroundColor: "rgba(0,0,0,0.85)",
border: "none",
borderRadius: "8px",
fontSize: 12,
}}
formatter={yFormatter}
labelStyle={{ color: "#fff" }}
/>
<Line
dataKey="bytesIn"
dot={false}
name="入站流量"
stroke="#10b981"
strokeWidth={2}
type="monotone"
/>
<Line
dataKey="bytesOut"
dot={false}
name="出站流量"
stroke="#ef4444"
strokeWidth={2}
type="monotone"
/>
<Line dataKey="bytesIn" dot={false} name="入站流量" stroke="#10b981" strokeWidth={2} type="monotone" />
<Line dataKey="bytesOut" dot={false} name="出站流量" stroke="#ef4444" strokeWidth={2} type="monotone" />
</LineChart>
</ResponsiveContainer>
</div>
@@ -333,6 +454,7 @@ function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
if (!hopsStr) return null;
let hops: TunnelQualityHopApiItem[] = [];
try {
hops = JSON.parse(hopsStr);
} catch {
@@ -353,28 +475,49 @@ function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
<div className="flex items-center overflow-x-auto pb-2 py-2">
{hops.map((hop, index) => {
const hasError = hop.latency < 0 || hop.loss > 0;
const colorClass = hop.latency < 0 ? "text-danger" : (hop.loss > 0 ? "text-warning" : "text-success");
const colorClass =
hop.latency < 0
? "text-danger"
: hop.loss > 0
? "text-warning"
: "text-success";
const borderColor = hasError ? "border-danger" : "";
return (
<React.Fragment key={index}>
{index === 0 && (
<Chip size="sm" variant="flat" className="shrink-0 font-mono shadow-sm">
<Chip
className="shrink-0 font-mono shadow-sm"
size="sm"
variant="flat"
>
{hop.fromNodeName}
</Chip>
)}
<div className="flex flex-col items-center justify-center min-w-[70px] mx-1 shrink-0 relative">
<span className={`text-[10px] font-mono leading-none mb-1 ${colorClass}`}>
<span
className={`text-[10px] font-mono leading-none mb-1 ${colorClass}`}
>
{hop.latency >= 0 ? `${hop.latency.toFixed(0)}ms` : "超时"}
</span>
<div className={`h-[2px] w-full relative flex items-center justify-end bg-default-200 ${hop.latency < 0 ? "!bg-danger" : ""}`}>
<ArrowRight className={`w-3.5 h-3.5 absolute -right-2 ${colorClass} bg-background rounded-full p-[1px] z-10`} />
<div
className={`h-[2px] w-full relative flex items-center justify-end bg-default-200 ${hop.latency < 0 ? "!bg-danger" : ""}`}
>
<ArrowRight
className={`w-3.5 h-3.5 absolute -right-2 ${colorClass} bg-background rounded-full p-[1px] z-10`}
/>
</div>
<span className={`text-[10px] font-mono leading-none mt-1.5 ${hop.loss > 0 ? "text-warning" : "text-default-400"}`}>
<span
className={`text-[10px] font-mono leading-none mt-1.5 ${hop.loss > 0 ? "text-warning" : "text-default-400"}`}
>
{hop.loss.toFixed(0)}% 丢包
</span>
</div>
<Chip size="sm" variant="flat" className={`shrink-0 font-mono shadow-sm ${borderColor}`}>
<Chip
className={`shrink-0 font-mono shadow-sm ${borderColor}`}
size="sm"
variant="flat"
>
{hop.toNodeName}
</Chip>
</React.Fragment>
@@ -386,15 +529,21 @@ function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
);
}
export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps) {
export function TunnelMonitorView({
viewMode = "grid",
}: TunnelMonitorViewProps) {
const [tunnels, setTunnels] = useState<MonitorTunnelApiItem[]>([]);
const [tunnelsLoading, setTunnelsLoading] = useState(false);
const [tunnelsError, setTunnelsError] = useState<string | null>(null);
const [accessDenied, setAccessDenied] = useState<string | null>(null);
// Quality data from backend periodic probing (latest per tunnel)
const [qualityMap, setQualityMap] = useState<Record<number, TunnelQualityApiItem>>({});
const [qualityHistoryMap, setQualityHistoryMap] = useState<Record<number, TunnelQualityApiItem[]>>({});
const [qualityMap, setQualityMap] = useState<
Record<number, TunnelQualityApiItem>
>({});
const [qualityHistoryMap, setQualityHistoryMap] = useState<
Record<number, TunnelQualityApiItem[]>
>({});
const initialHistoryFetched = useRef(false);
const [qualityLoading, setQualityLoading] = useState(false);
const qualityTimerRef = useRef<number | null>(null);
@@ -405,20 +554,27 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const [detailTunnelId, setDetailTunnelId] = useState<number | null>(null);
// Quality history for chart (mirrors service monitor results)
const [qualityHistory, setQualityHistory] = useState<TunnelQualityApiItem[]>([]);
const [qualityHistory, setQualityHistory] = useState<TunnelQualityApiItem[]>(
[],
);
const [qualityHistoryLoading, setQualityHistoryLoading] = useState(false);
const [qualityHistoryError, setQualityHistoryError] = useState<string | null>(null);
const [qualityHistoryError, setQualityHistoryError] = useState<string | null>(
null,
);
const [qualityRangeMs, setQualityRangeMs] = useState(60 * 60 * 1000);
// Tunnel traffic metrics for chart
const [tunnelMetrics, setTunnelMetrics] = useState<TunnelMetricApiItem[]>([]);
const [tunnelMetricsLoading, setTunnelMetricsLoading] = useState(false);
const [tunnelMetricsError, setTunnelMetricsError] = useState<string | null>(null);
const [tunnelMetricsError, setTunnelMetricsError] = useState<string | null>(
null,
);
const [tunnelRangeMs, setTunnelRangeMs] = useState(60 * 60 * 1000);
// --- Load tunnel list ---
const loadTunnels = useCallback(async (options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
if (!silent) setTunnelsLoading(true);
try {
const response = await getMonitorTunnels();
@@ -427,12 +583,14 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
setAccessDenied(null);
setTunnelsError(null);
setTunnels(response.data);
return;
}
if (response.code === 403) {
setAccessDenied(response.msg || "暂无监控权限,请联系管理员授权");
setTunnelsError(null);
setTunnels([]);
return;
}
setTunnelsError(response.msg || "加载隧道列表失败");
@@ -452,6 +610,7 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const response = await getConfigByName(
MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY,
);
setMonitorTunnelQualityEnabled(
typeof response.data?.value === "string"
? response.data.value === "true"
@@ -477,7 +636,9 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
useEffect(() => {
const handleMonitorTunnelQualityEnabledChanged = (event: Event) => {
const enabled = (event as CustomEvent<{ enabled?: boolean }>).detail?.enabled;
const enabled = (event as CustomEvent<{ enabled?: boolean }>).detail
?.enabled;
if (typeof enabled === "boolean") {
setMonitorTunnelQualityEnabled(enabled);
if (!enabled) {
@@ -510,31 +671,44 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const start = end - 5 * 60 * 1000;
const promises = tunnels.map((t) =>
getMonitorTunnelQualityHistory(t.id, start, end).then(res => {
if (res.code === 0 && res.data) {
return { id: t.id, data: res.data };
}
return { id: t.id, data: [] };
}).catch(() => ({ id: t.id, data: [] }))
getMonitorTunnelQualityHistory(t.id, start, end)
.then((res) => {
if (res.code === 0 && res.data) {
return { id: t.id, data: res.data };
}
return { id: t.id, data: [] };
})
.catch(() => ({ id: t.id, data: [] })),
);
const results = await Promise.all(promises);
setQualityHistoryMap(prev => {
setQualityHistoryMap((prev) => {
const nextMap = { ...prev };
results.forEach(r => {
results.forEach((r) => {
const existing = nextMap[r.id] || [];
const merged = [...r.data, ...existing].sort((a,b) => a.timestamp - b.timestamp);
const merged = [...r.data, ...existing].sort(
(a, b) => a.timestamp - b.timestamp,
);
const unique: TunnelQualityApiItem[] = [];
for (const item of merged) {
if (unique.length === 0 || unique[unique.length-1].timestamp !== item.timestamp) {
unique.push(item);
}
if (
unique.length === 0 ||
unique[unique.length - 1].timestamp !== item.timestamp
) {
unique.push(item);
}
}
nextMap[r.id] = unique.slice(-UPTIME_KUMA_BARS_COUNT);
});
return nextMap;
});
};
void fetchHistory();
}
}, [tunnels]);
@@ -542,27 +716,36 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
// --- Load quality snapshots (auto-polling every 10s) ---
const loadQuality = useCallback(async (options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
if (!silent) setQualityLoading(true);
try {
const response = await getMonitorTunnelQuality();
if (response.code === 0 && Array.isArray(response.data)) {
const map: Record<number, TunnelQualityApiItem> = {};
for (const q of response.data) {
map[q.tunnelId] = q;
}
setQualityMap(map);
// Accumulate locally
setQualityHistoryMap(prev => {
const nextMap = { ...prev };
for (const q of response.data) {
const currentArr = nextMap[q.tunnelId] || [];
const lastItem = currentArr.length > 0 ? currentArr[currentArr.length - 1] : null;
if (!lastItem || lastItem.timestamp !== q.timestamp) {
nextMap[q.tunnelId] = [...currentArr, q].slice(-UPTIME_KUMA_BARS_COUNT);
}
}
return nextMap;
setQualityHistoryMap((prev) => {
const nextMap = { ...prev };
for (const q of response.data) {
const currentArr = nextMap[q.tunnelId] || [];
const lastItem =
currentArr.length > 0 ? currentArr[currentArr.length - 1] : null;
if (!lastItem || lastItem.timestamp !== q.timestamp) {
nextMap[q.tunnelId] = [...currentArr, q].slice(
-UPTIME_KUMA_BARS_COUNT,
);
}
}
return nextMap;
});
}
} catch {
@@ -586,6 +769,7 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
window.clearInterval(qualityTimerRef.current);
qualityTimerRef.current = null;
}
return;
}
@@ -605,19 +789,26 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const loadQualityHistory = useCallback(
async (tunnelId: number, options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
if (!silent) setQualityHistoryLoading(true);
try {
const end = Date.now();
const start = end - qualityRangeMs;
const response = await getMonitorTunnelQualityHistory(tunnelId, start, end);
const response = await getMonitorTunnelQualityHistory(
tunnelId,
start,
end,
);
if (response.code === 0 && Array.isArray(response.data)) {
setQualityHistoryError(null);
setQualityHistory(response.data);
return;
}
if (response.code === 403) {
setAccessDenied(response.msg || "暂无监控权限");
return;
}
setQualityHistoryError(response.msg || "加载质量历史失败");
@@ -637,6 +828,7 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const loadTunnelMetrics = useCallback(
async (tunnelId: number, options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
if (!silent) setTunnelMetricsLoading(true);
try {
const end = Date.now();
@@ -648,7 +840,9 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const ordered = [...response.data].sort(
(a, b) => a.timestamp - b.timestamp,
);
setTunnelMetrics(ordered);
return;
}
setTunnelMetricsError(response.msg || "加载流量数据失败");
@@ -695,29 +889,32 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
// Memoize chart data so React.memo sub-components see stable references
const qualityChartData = useMemo(
() => qualityHistory.map((q) => ({
time: formatTimestamp(q.timestamp, qualityRangeMs),
entryToExit: q.entryToExitLatency >= 0 ? q.entryToExitLatency : null,
exitToBing: q.exitToBingLatency >= 0 ? q.exitToBingLatency : null,
entryToExitLoss: q.entryToExitLoss,
exitToBingLoss: q.exitToBingLoss,
})),
() =>
qualityHistory.map((q) => ({
time: formatTimestamp(q.timestamp, qualityRangeMs),
entryToExit: q.entryToExitLatency >= 0 ? q.entryToExitLatency : null,
exitToBing: q.exitToBingLatency >= 0 ? q.exitToBingLatency : null,
entryToExitLoss: q.entryToExitLoss,
exitToBingLoss: q.exitToBingLoss,
})),
[qualityHistory, qualityRangeMs],
);
const tunnelChartData = useMemo(
() => tunnelMetrics.map((m) => ({
time: formatTimestamp(m.timestamp, tunnelRangeMs),
bytesIn: m.bytesIn,
bytesOut: m.bytesOut,
connections: m.connections,
})),
() =>
tunnelMetrics.map((m) => ({
time: formatTimestamp(m.timestamp, tunnelRangeMs),
bytesIn: m.bytesIn,
bytesOut: m.bytesOut,
connections: m.connections,
})),
[tunnelMetrics, tunnelRangeMs],
);
const detailTunnel = detailTunnelId != null
? tunnels.find((t) => t.id === detailTunnelId)
: null;
const detailTunnel =
detailTunnelId != null
? tunnels.find((t) => t.id === detailTunnelId)
: null;
// Aggregate stats
const tunnelStats = useMemo(() => {
@@ -729,9 +926,11 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
// Last quality update timestamp
const lastQualityUpdate = useMemo(() => {
let latest = 0;
for (const q of Object.values(qualityMap)) {
if (q.timestamp > latest) latest = q.timestamp;
}
return latest > 0 ? new Date(latest).toLocaleTimeString("zh-CN") : null;
}, [qualityMap]);
@@ -759,18 +958,28 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
<div className="space-y-6">
{/* Header */}
<div className="flex items-center gap-3 flex-wrap">
<Button size="sm" variant="flat" onPress={() => {
setDetailTunnelId(null);
setQualityHistory([]);
setTunnelMetrics([]);
}}>
<Button
size="sm"
variant="flat"
onPress={() => {
setDetailTunnelId(null);
setQualityHistory([]);
setTunnelMetrics([]);
}}
>
<ArrowLeft className="w-4 h-4 mr-1" />
返回隧道列表
</Button>
<div className="flex items-center gap-2">
<ArrowRightLeft className={`w-5 h-5 ${detailTunnel.status === 1 ? "text-success" : "text-default-400"}`} />
<ArrowRightLeft
className={`w-5 h-5 ${detailTunnel.status === 1 ? "text-success" : "text-default-400"}`}
/>
<h3 className="text-lg font-semibold">{detailTunnel.name}</h3>
<Chip size="sm" color={detailTunnel.status === 1 ? "success" : "danger"} variant="flat">
<Chip
color={detailTunnel.status === 1 ? "success" : "danger"}
size="sm"
variant="flat"
>
{detailTunnel.status === 1 ? "启用" : "禁用"}
</Chip>
</div>
@@ -784,7 +993,10 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
<Zap className="w-3 h-3" />
入口 → 出口 延迟
</span>
<LatencyDisplay value={quality?.entryToExitLatency} loading={qualityLoading} />
<LatencyDisplay
loading={qualityLoading}
value={quality?.entryToExitLatency}
/>
</CardBody>
</Card>
<Card className="border border-divider/60 shadow-sm hover:shadow-md transition-shadow bg-gradient-to-br from-background to-default-50/50">
@@ -793,22 +1005,37 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
<Globe className="w-3 h-3" />
出口 → Bing 延迟
</span>
<LatencyDisplay value={quality?.exitToBingLatency} loading={qualityLoading} />
<LatencyDisplay
loading={qualityLoading}
value={quality?.exitToBingLatency}
/>
</CardBody>
</Card>
<Card className="border border-divider/60 shadow-sm hover:shadow-md transition-shadow bg-gradient-to-br from-background to-default-50/50">
<CardBody className="py-3 px-4 flex flex-col items-center justify-center min-h-[5rem]">
<span className="text-[11px] text-default-500 mb-1.5">入口 → 出口 丢包</span>
<span className={`text-sm font-semibold font-mono ${(quality?.entryToExitLoss ?? 0) > 0 ? "text-warning" : ""}`}>
{quality?.entryToExitLoss !== undefined ? `${quality.entryToExitLoss.toFixed(1)}%` : "-"}
<span className="text-[11px] text-default-500 mb-1.5">
入口 → 出口 丢包
</span>
<span
className={`text-sm font-semibold font-mono ${(quality?.entryToExitLoss ?? 0) > 0 ? "text-warning" : ""}`}
>
{quality?.entryToExitLoss !== undefined
? `${quality.entryToExitLoss.toFixed(1)}%`
: "-"}
</span>
</CardBody>
</Card>
<Card className="border border-divider/60 shadow-sm hover:shadow-md transition-shadow bg-gradient-to-br from-background to-default-50/50">
<CardBody className="py-3 px-4 flex flex-col items-center justify-center min-h-[5rem]">
<span className="text-[11px] text-default-500 mb-1.5">出口 → Bing 丢包</span>
<span className={`text-sm font-semibold font-mono ${(quality?.exitToBingLoss ?? 0) > 0 ? "text-warning" : ""}`}>
{quality?.exitToBingLoss !== undefined ? `${quality.exitToBingLoss.toFixed(1)}%` : "-"}
<span className="text-[11px] text-default-500 mb-1.5">
出口 → Bing 丢包
</span>
<span
className={`text-sm font-semibold font-mono ${(quality?.exitToBingLoss ?? 0) > 0 ? "text-warning" : ""}`}
>
{quality?.exitToBingLoss !== undefined
? `${quality.exitToBingLoss.toFixed(1)}%`
: "-"}
</span>
</CardBody>
</Card>
@@ -829,7 +1056,8 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
)}
{quality?.timestamp && (
<span className="text-default-400">
· 最近更新: {new Date(quality.timestamp).toLocaleTimeString("zh-CN")}
· 最近更新:{" "}
{new Date(quality.timestamp).toLocaleTimeString("zh-CN")}
</span>
)}
{quality?.errorMessage && (
@@ -844,23 +1072,23 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
{/* ====== Quality History Chart — isolated with React.memo ====== */}
<QualityChartCard
rangeMs={qualityRangeMs}
onRangeChange={setQualityRangeMs}
loading={qualityHistoryLoading}
error={qualityHistoryError}
data={qualityChartData}
error={qualityHistoryError}
loading={qualityHistoryLoading}
rangeMs={qualityRangeMs}
tunnelId={detailTunnelId}
onRangeChange={setQualityRangeMs}
onRefresh={loadQualityHistory}
/>
{/* ====== Traffic Chart — isolated with React.memo ====== */}
<TrafficChartCard
rangeMs={tunnelRangeMs}
onRangeChange={setTunnelRangeMs}
loading={tunnelMetricsLoading}
error={tunnelMetricsError}
data={tunnelChartData}
error={tunnelMetricsError}
loading={tunnelMetricsLoading}
rangeMs={tunnelRangeMs}
tunnelId={detailTunnelId}
onRangeChange={setTunnelRangeMs}
onRefresh={loadTunnelMetrics}
/>
</div>
@@ -871,7 +1099,9 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
return (
<div className="space-y-6">
<div className="flex flex-wrap items-center gap-3 mb-1">
<Chip color="primary" size="sm" variant="flat">隧道 {tunnelStats.enabled}/{tunnelStats.total}</Chip>
<Chip color="primary" size="sm" variant="flat">
隧道 {tunnelStats.enabled}/{tunnelStats.total}
</Chip>
{lastQualityUpdate ? (
<div className="flex items-center gap-1.5 text-xs text-default-500">
{monitorTunnelQualityEnabled ? (
@@ -893,7 +1123,12 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
</div>
) : null}
<div className="ml-auto">
<Button isLoading={tunnelsLoading} size="sm" variant="flat" onPress={() => loadTunnels()}>
<Button
isLoading={tunnelsLoading}
size="sm"
variant="flat"
onPress={() => loadTunnels()}
>
<RefreshCw className="w-4 h-4 mr-1" />
刷新
</Button>
@@ -920,21 +1155,33 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
className="group relative overflow-hidden shadow-sm border border-divider dark:border-default-100 hover:-translate-y-1 hover:shadow-lg transition-all duration-300 h-full flex flex-col cursor-pointer bg-background"
onClick={() => setDetailTunnelId(tunnel.id)}
>
<div className={`absolute top-0 left-0 right-0 h-1 ${isEnabled ? "bg-success" : "bg-danger"}`} />
<div className={`absolute -right-8 -top-8 w-24 h-24 rounded-full blur-2xl opacity-10 transition-opacity group-hover:opacity-20 ${isEnabled ? "bg-success" : "bg-danger"}`} />
<div
className={`absolute top-0 left-0 right-0 h-1 ${isEnabled ? "bg-success" : "bg-danger"}`}
/>
<div
className={`absolute -right-8 -top-8 w-24 h-24 rounded-full blur-2xl opacity-10 transition-opacity group-hover:opacity-20 ${isEnabled ? "bg-success" : "bg-danger"}`}
/>
<CardHeader className="pb-2 pt-5 px-5 flex flex-row justify-between items-start gap-4">
<div className="flex items-center gap-3 min-w-0">
<div className="relative flex-shrink-0">
<div className="w-10 h-10 rounded-xl bg-default-100 dark:bg-default-50/10 flex items-center justify-center border border-divider">
<ArrowRightLeft className={`w-5 h-5 ${isEnabled ? "text-success" : "text-danger"}`} />
<ArrowRightLeft
className={`w-5 h-5 ${isEnabled ? "text-success" : "text-danger"}`}
/>
</div>
<span className={`absolute -bottom-0.5 -right-0.5 w-3 h-3 rounded-full border-2 border-background ${isEnabled ? "bg-success" : "bg-danger"}`} />
<span
className={`absolute -bottom-0.5 -right-0.5 w-3 h-3 rounded-full border-2 border-background ${isEnabled ? "bg-success" : "bg-danger"}`}
/>
</div>
<div className="flex flex-col min-w-0">
<h3 className="font-semibold text-foreground text-sm truncate">{tunnel.name}</h3>
<h3 className="font-semibold text-foreground text-sm truncate">
{tunnel.name}
</h3>
<div className="flex items-center gap-1.5 text-[11px] text-default-500 mt-0.5">
<span className="font-mono">{isEnabled ? "启用" : "禁用"}</span>
<span className="font-mono">
{isEnabled ? "启用" : "禁用"}
</span>
</div>
</div>
</div>
@@ -947,20 +1194,30 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
<Zap className="w-3 h-3" />
入口→出口
</div>
<UptimeHistoryBar history={qualityHistoryMap[tunnel.id]} type="entryToExit" latestValue={quality?.entryToExitLatency} />
<UptimeHistoryBar
history={qualityHistoryMap[tunnel.id]}
latestValue={quality?.entryToExitLatency}
type="entryToExit"
/>
</div>
<div className="space-y-1">
<div className="text-[10px] text-default-500 flex items-center gap-1">
<Globe className="w-3 h-3" />
出口→Bing
</div>
<UptimeHistoryBar history={qualityHistoryMap[tunnel.id]} type="exitToBing" latestValue={quality?.exitToBingLatency} />
<UptimeHistoryBar
history={qualityHistoryMap[tunnel.id]}
latestValue={quality?.exitToBingLatency}
type="exitToBing"
/>
</div>
</div>
<div className="flex justify-between items-center pt-2 border-t border-divider/50">
{quality?.errorMessage ? (
<span className="text-[11px] text-danger truncate">{quality.errorMessage}</span>
<span className="text-[11px] text-danger truncate">
{quality.errorMessage}
</span>
) : quality?.timestamp ? (
<span className="text-[11px] text-default-500 flex items-center gap-1">
{monitorTunnelQualityEnabled ? (
@@ -968,11 +1225,15 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
) : (
<WifiOff className="w-3 h-3 text-warning" />
)}
{new Date(quality.timestamp).toLocaleTimeString("zh-CN")}
{new Date(quality.timestamp).toLocaleTimeString(
"zh-CN",
)}
</span>
) : (
<span className="text-[11px] text-default-400">
{monitorTunnelQualityEnabled ? "等待探测..." : "实时检测已关闭"}
{monitorTunnelQualityEnabled
? "等待探测..."
: "实时检测已关闭"}
</span>
)}
</div>
@@ -987,7 +1248,8 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
aria-label="隧道列表"
className="overflow-x-auto min-w-full"
classNames={{
wrapper: "bg-transparent p-0 shadow-none border-none overflow-hidden rounded-2xl",
wrapper:
"bg-transparent p-0 shadow-none border-none overflow-auto rounded-2xl",
th: "bg-transparent text-default-600 font-semibold text-sm border-b border-white/20 dark:border-white/10 py-3 uppercase tracking-wider first:rounded-tl-[24px] last:rounded-tr-[24px]",
td: "py-3 border-b border-divider/50 group-data-[last=true]:border-b-0",
tr: "hover:bg-white/40 dark:hover:bg-white/10 transition-colors",
@@ -1006,7 +1268,11 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
const isEnabled = tunnel.status === 1;
return (
<TableRow key={tunnel.id} className="cursor-pointer" onClick={() => setDetailTunnelId(tunnel.id)}>
<TableRow
key={tunnel.id}
className="cursor-pointer"
onClick={() => setDetailTunnelId(tunnel.id)}
>
<TableCell>
<div className="flex items-center gap-1.5">
{isEnabled ? (
@@ -1017,13 +1283,23 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
</div>
</TableCell>
<TableCell>
<span className="font-semibold text-sm whitespace-nowrap">{tunnel.name}</span>
<span className="font-semibold text-sm whitespace-nowrap">
{tunnel.name}
</span>
</TableCell>
<TableCell>
<UptimeHistoryBar history={qualityHistoryMap[tunnel.id]} type="entryToExit" latestValue={quality?.entryToExitLatency} />
<UptimeHistoryBar
history={qualityHistoryMap[tunnel.id]}
latestValue={quality?.entryToExitLatency}
type="entryToExit"
/>
</TableCell>
<TableCell>
<UptimeHistoryBar history={qualityHistoryMap[tunnel.id]} type="exitToBing" latestValue={quality?.exitToBingLatency} />
<UptimeHistoryBar
history={qualityHistoryMap[tunnel.id]}
latestValue={quality?.exitToBingLatency}
type="exitToBing"
/>
</TableCell>
<TableCell>
{quality?.timestamp ? (
@@ -1033,7 +1309,9 @@ export function TunnelMonitorView({ viewMode = "grid" }: TunnelMonitorViewProps)
) : (
<WifiOff className="w-3.5 h-3.5 text-warning" />
)}
{new Date(quality.timestamp).toLocaleTimeString("zh-CN")}
{new Date(quality.timestamp).toLocaleTimeString(
"zh-CN",
)}
</span>
) : (
<span className="text-xs text-default-400">

Some files were not shown because too many files have changed in this diff Show More