mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
Compare commits
65 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 25dfb84324 | |||
| 4ebd6703fe | |||
| 1f53a39784 | |||
| 5ebd4c2a91 | |||
| 6c93d829c6 | |||
| 5d22d4cb06 | |||
| e5cd5af550 | |||
| cdcdfd8ff0 | |||
| 791773fd62 | |||
| 13764b4615 | |||
| 4c882d907b | |||
| 6033e39466 | |||
| 0f3242bf11 | |||
| d97d91801d | |||
| 727ef56c67 | |||
| a40150b136 | |||
| a923ec4785 | |||
| 42c5492c1d | |||
| 869d726b7a | |||
| 55a931510b | |||
| a259dd83b2 | |||
| cc0b8de2e1 | |||
| 58ef260755 | |||
| 521fe79b15 | |||
| 90012725cc | |||
| 615d9e67eb | |||
| a131b70613 | |||
| cbed4eab23 | |||
| c2745dcd56 | |||
| ad4109594a | |||
| efc8c75dcb | |||
| e5acc49186 | |||
| d98377a297 | |||
| 950e9a9ba8 | |||
| 3f3159aafd | |||
| 5e8d0682c0 | |||
| 3c0e833cfc | |||
| 60311d3e47 | |||
| b4192c9e94 | |||
| f05e9480ee | |||
| 9861b44107 | |||
| d8144821e6 | |||
| 023be27287 | |||
| e8d5687419 | |||
| 0fbe570597 | |||
| 4c52d7fec2 | |||
| 3373e5ade9 | |||
| bd27b94909 | |||
| a5a500bc0f | |||
| edfe2a2372 | |||
| 7a9ba8bd81 | |||
| 46394388b1 | |||
| dec337d46b | |||
| 9e8d27d98e | |||
| a2000e4d98 | |||
| 2b76a9f0be | |||
| 54d7dfb7c9 | |||
| 9f19d5fe15 | |||
| 2ca3849917 | |||
| 58d2e89147 | |||
| 87a1a34ad5 | |||
| a625884d61 | |||
| 799bb66fe5 | |||
| 3f374df724 | |||
| 9b923a2d0b |
@@ -1,6 +1,6 @@
|
||||
# FLVX
|
||||
|
||||
> **联系我们**: [Telegram群组](https://t.me/flvxpanel)
|
||||
> **联系我们**: [Telegram群组](https://t.me/flvxchannel)
|
||||
|
||||
|
||||
## 特性
|
||||
@@ -184,7 +184,6 @@ This fork (FLVX) is no longer a light patch on top of the upstream project. It h
|
||||
|
||||
| 网络 | 地址 |
|
||||
|------------|----------------------------------------------------------------------|
|
||||
| BNB(BEP20) | `0xa608708fdc6279a2433fd4b82f0b72b8cbe97ed5` |
|
||||
| TRC20 | `TM8VYdU3s3gSX5PC8swjAJrAzZFCHKqG2k` |
|
||||
| Aptos | `0x49427bfcba1006a346447430689b2307ac156316bb34850d1d3029ff9d118da5` |
|
||||
| polygon | `0xa608708fdc6279a2433fd4b82f0b72b8cbe97ed5` |
|
||||
| BNB(BEP20) | `0x271327ce49140e670eA0F772d9886BF90E9022Ee` |
|
||||
| TRC20 | `TARxZWggaxFqYgxGVBxPkyykgYKNmGndmE` |
|
||||
| polygon | `0x271327ce49140e670eA0F772d9886BF90E9022Ee` |
|
||||
|
||||
@@ -15,10 +15,15 @@ services:
|
||||
JWT_SECRET: ${JWT_SECRET}
|
||||
SERVER_ADDR: :6365
|
||||
TZ: Asia/Shanghai
|
||||
FLUX_VERSION: ${FLUX_VERSION:-dev}
|
||||
PANEL_DEPLOY_DIR: /opt/flvx-panel
|
||||
PANEL_BACKEND_CONTAINER: flux-panel-backend
|
||||
ports:
|
||||
- "${BACKEND_PORT}:6365"
|
||||
volumes:
|
||||
- sqlite_data:/app/data
|
||||
- /var/run/docker.sock:/var/run/docker.sock
|
||||
- ./:/opt/flvx-panel
|
||||
networks:
|
||||
- gost-network
|
||||
stop_grace_period: 30s
|
||||
|
||||
@@ -15,10 +15,15 @@ services:
|
||||
JWT_SECRET: ${JWT_SECRET}
|
||||
SERVER_ADDR: :6365
|
||||
TZ: Asia/Shanghai
|
||||
FLUX_VERSION: ${FLUX_VERSION:-dev}
|
||||
PANEL_DEPLOY_DIR: /opt/flvx-panel
|
||||
PANEL_BACKEND_CONTAINER: flux-panel-backend
|
||||
ports:
|
||||
- "${BACKEND_PORT}:6365"
|
||||
volumes:
|
||||
- sqlite_data:/app/data
|
||||
- /var/run/docker.sock:/var/run/docker.sock
|
||||
- ./:/opt/flvx-panel
|
||||
networks:
|
||||
- gost-network
|
||||
stop_grace_period: 30s
|
||||
|
||||
@@ -0,0 +1,171 @@
|
||||
# Allow Local Remote Address Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Add a global settings toggle that allows non-admin forward rules to target local/private addresses when explicitly enabled.
|
||||
|
||||
**Architecture:** Keep the existing remote-address safety validator as the default path for non-admin rule changes, but gate its use behind a single backend config lookup in forward create/update handlers. Surface the toggle through the existing `vite_config` settings page and prove behavior with backend contract tests first.
|
||||
|
||||
**Tech Stack:** Go `net/http` + GORM backend, React + TypeScript frontend settings page, Go contract tests.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Backend Contract Coverage
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/tests/contract/forward_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing tests**
|
||||
|
||||
Add contract tests that prove the desired behavior:
|
||||
|
||||
```go
|
||||
t.Run("local remote address is rejected when toggle is off", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "deny-local-remote",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, _ := json.Marshal(createPayload)
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
|
||||
var out response.R
|
||||
_ = json.NewDecoder(createRes.Body).Decode(&out)
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected local remote address to be rejected when toggle is off")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("local remote address is allowed when toggle is on", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, "allow_local_remote_addr", "1", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("enable allow_local_remote_addr: %v", err)
|
||||
}
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "allow-local-remote",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, _ := json.Marshal(createPayload)
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
})
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run tests to verify they fail**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: FAIL because backend still rejects local/private addresses unconditionally.
|
||||
|
||||
- [ ] **Step 3: Commit**
|
||||
|
||||
Do not commit yet; combine with Task 2 after implementation passes.
|
||||
|
||||
### Task 2: Backend Toggle Implementation
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
|
||||
- [ ] **Step 1: Add a tiny config helper**
|
||||
|
||||
Add a helper near other handler helpers:
|
||||
|
||||
```go
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(cfg.Value) == "1"
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Gate create/update validation behind the helper**
|
||||
|
||||
Replace the unconditional checks with:
|
||||
|
||||
```go
|
||||
if !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Run contract tests to verify they pass**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 4: Run full backend tests**
|
||||
|
||||
Run: `go test ./...`
|
||||
Expected: PASS
|
||||
|
||||
### Task 3: Settings Page Toggle
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
- [ ] **Step 1: Add the config item to the settings schema**
|
||||
|
||||
Add a switch-style item for `allow_local_remote_addr` with warning copy about reduced safety.
|
||||
|
||||
- [ ] **Step 2: Ensure the key is included in config loading/saving paths**
|
||||
|
||||
Add `allow_local_remote_addr` anywhere the page enumerates config keys or groups persisted config values.
|
||||
|
||||
- [ ] **Step 3: Run frontend build**
|
||||
|
||||
Run: `pnpm run build`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 4: Run frontend lint**
|
||||
|
||||
Run: `pnpm run lint`
|
||||
Expected: 0 errors; existing warnings may remain.
|
||||
|
||||
### Task 4: Final Verification
|
||||
|
||||
**Files:**
|
||||
- Verify only
|
||||
|
||||
- [ ] **Step 1: Re-run backend contracts for the toggle**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 2: Re-run full backend tests**
|
||||
|
||||
Run: `go test ./...`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 3: Re-run frontend build/lint**
|
||||
|
||||
Run: `pnpm run build && pnpm run lint`
|
||||
Expected: Build passes, lint has no errors.
|
||||
|
||||
- [ ] **Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/mutations.go go-backend/tests/contract/forward_contract_test.go vite-frontend/src/pages/config.tsx docs/superpowers/specs/2026-04-26-allow-local-remote-addr-design.md docs/superpowers/plans/2026-04-26-allow-local-remote-addr.md
|
||||
git commit -m "feat: add allow-local-remote-address toggle"
|
||||
```
|
||||
@@ -0,0 +1,849 @@
|
||||
# flow/upload Batch Optimization Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Reduce `POST /flow/upload` database pressure by converting the hot path from per-item queries and per-item transactions to per-request aggregation, batched metadata reads, and batched writes, while preserving immediate quota disable / forward pause behavior inside the same upload.
|
||||
|
||||
**Architecture:** Parse one upload into a batch object in the handler layer, fetch one shared `forward+tunnel` metadata map, then reuse that map for flow accounting and tunnel metric aggregation. Replace `AddFlow` and `AddUserQuotaUsage` per-item transactions with one batched flow transaction and one batched quota transaction; run policy enforcement, orphan cleanup, and peer-share flow handling once per affected target instead of once per item.
|
||||
|
||||
**Tech Stack:** Go, net/http, GORM, SQLite/PostgreSQL, existing backend contract tests.
|
||||
|
||||
---
|
||||
|
||||
## File Map
|
||||
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch.go`
|
||||
Responsibility: request-scoped parsing, aggregation, and application of one `/flow/upload` batch.
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch_test.go`
|
||||
Responsibility: unit coverage for batch aggregation semantics.
|
||||
- Create: `go-backend/internal/store/repo/repository_flow_batch_test.go`
|
||||
Responsibility: unit coverage for batched flow and quota persistence.
|
||||
- Create: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
Responsibility: contract coverage that repeated items still accumulate correctly and still disable quota immediately.
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
Responsibility: switch `/flow/upload` entrypoint to the new batch pipeline.
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_metrics_ingestion.go`
|
||||
Responsibility: accept pre-aggregated forward deltas plus shared forward metadata instead of reparsing the raw items.
|
||||
- Modify: `go-backend/internal/store/repo/repository.go`
|
||||
Responsibility: add batched flow persistence primitives near the existing flow update code.
|
||||
- Modify: `go-backend/internal/store/repo/repository_flow.go`
|
||||
Responsibility: add shared flow-upload metadata query helpers.
|
||||
- Modify: `go-backend/internal/store/repo/repository_user_quota.go`
|
||||
Responsibility: add batched quota usage persistence that still returns normalized quota views for immediate enforcement.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Add Failing Tests For Batched flow/upload Semantics
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch_test.go`
|
||||
- Create: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing handler unit test**
|
||||
|
||||
Create `go-backend/internal/http/handler/flow_upload_batch_test.go` with a unit test that locks in the new aggregation contract.
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets(t *testing.T) {
|
||||
h := &Handler{}
|
||||
metas := map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {
|
||||
ForwardID: 20,
|
||||
TunnelID: 1,
|
||||
TrafficRatio: 2,
|
||||
TunnelFlow: 3,
|
||||
},
|
||||
}
|
||||
|
||||
batch := h.buildFlowUploadBatch([]flowItem{
|
||||
{N: "20_2_10", U: 70, D: 50},
|
||||
{N: "20_2_10_tcp", U: 40, D: 30},
|
||||
{N: "99_2_10", U: 12, D: 8},
|
||||
{N: "fed_svc_17", U: 9, D: 1},
|
||||
}, metas)
|
||||
|
||||
if len(batch.flowDeltas) != 1 {
|
||||
t.Fatalf("expected 1 flow delta, got %d", len(batch.flowDeltas))
|
||||
}
|
||||
delta := batch.flowDeltas[0]
|
||||
if delta.ForwardID != 20 || delta.UserID != 2 || delta.UserTunnelID != 10 {
|
||||
t.Fatalf("unexpected flow delta identity: %#v", delta)
|
||||
}
|
||||
if delta.InFlow != 480 || delta.OutFlow != 660 {
|
||||
t.Fatalf("expected scaled flow in=480 out=660, got in=%d out=%d", delta.InFlow, delta.OutFlow)
|
||||
}
|
||||
if batch.quotaUsage[2] != 1140 {
|
||||
t.Fatalf("expected quota usage 1140, got %d", batch.quotaUsage[2])
|
||||
}
|
||||
if len(batch.policyTargets) != 1 {
|
||||
t.Fatalf("expected 1 policy target, got %d", len(batch.policyTargets))
|
||||
}
|
||||
if batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("unexpected policy target: %#v", batch.policyTargets[0])
|
||||
}
|
||||
traffic := batch.forwardTraffic[20]
|
||||
if traffic.bytesIn != 80 || traffic.bytesOut != 110 {
|
||||
t.Fatalf("expected raw traffic in=80 out=110, got in=%d out=%d", traffic.bytesIn, traffic.bytesOut)
|
||||
}
|
||||
if _, ok := batch.orphanServices["99_2_10"]; !ok {
|
||||
t.Fatalf("expected orphan service cleanup target for 99_2_10")
|
||||
}
|
||||
if item, ok := batch.peerShareForwardItems["20_2_10"]; !ok || item.U != 110 || item.D != 80 {
|
||||
t.Fatalf("expected merged peer-share forward item, got %#v ok=%v", item, ok)
|
||||
}
|
||||
if item, ok := batch.peerShareRuntimeItems[17]; !ok || item.U != 9 || item.D != 1 {
|
||||
t.Fatalf("expected merged peer-share runtime item, got %#v ok=%v", item, ok)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the handler unit test to verify RED**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets -v
|
||||
```
|
||||
|
||||
Expected: FAIL because `FlowUploadForwardMeta`, `buildFlowUploadBatch`, and the new batch fields do not exist yet.
|
||||
|
||||
- [ ] **Step 3: Write the contract test that guards current behavior**
|
||||
|
||||
Create `go-backend/tests/contract/flow_upload_batch_contract_test.go` so the optimization cannot weaken same-request quota enforcement.
|
||||
|
||||
```go
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately(t *testing.T) {
|
||||
secret := "monitoring-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
|
||||
monthKey := int64(now.Year()*100 + int(now.Month()))
|
||||
const bytesPerGB = int64(1024 * 1024 * 1024)
|
||||
|
||||
node := &model.Node{Name: "node-1", Secret: "node-secret", ServerIP: "127.0.0.1", Port: "10000-10010", TCPListenAddr: "[::]", UDPListenAddr: "[::]", CreatedTime: nowMs, Status: 1}
|
||||
if err := repo.DB().Create(node).Error; err != nil {
|
||||
t.Fatalf("seed node: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'flow_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
tunnel := &model.Tunnel{Name: "tunnel-1", TrafficRatio: 1.0, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
|
||||
if err := repo.DB().Create(tunnel).Error; err != nil {
|
||||
t.Fatalf("seed tunnel: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`, tunnel.ID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
forward := &model.Forward{ID: 20, UserID: 2, UserName: "flow_user", Name: "forward-20", TunnelID: tunnel.ID, RemoteAddr: "1.1.1.1:80", Strategy: "fifo", CreatedTime: nowMs, UpdatedTime: nowMs, Status: 1}
|
||||
if err := repo.DB().Create(forward).Error; err != nil {
|
||||
t.Fatalf("seed forward: %v", err)
|
||||
}
|
||||
if err := repo.DB().Exec(`INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time) VALUES(2, 1, 0, ?, ?, ?, ?, 0, 0, '', ?, ?)`, bytesPerGB-100, bytesPerGB-100, dayKey, monthKey, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user_quota: %v", err)
|
||||
}
|
||||
|
||||
body, err := json.Marshal([]map[string]interface{}{
|
||||
{"n": "20_2_10", "u": 70, "d": 50},
|
||||
{"n": "20_2_10_tcp", "u": 40, "d": 30},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal body: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected status 200, got %d", res.Code)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT status FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected forward paused immediately, got status=%d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`); got != 1 {
|
||||
t.Fatalf("expected quota disabled flag=1, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 80 {
|
||||
t.Fatalf("expected forward in_flow=80, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 110 {
|
||||
t.Fatalf("expected forward out_flow=110, got %d", got)
|
||||
}
|
||||
metrics, err := repo.GetTunnelMetrics(tunnel.ID, 0, nowMs+60_000)
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 || metrics[0].BytesIn != 80 || metrics[0].BytesOut != 110 {
|
||||
t.Fatalf("expected one aggregated metric row, got %#v", metrics)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run the contract test to verify the same-request guard stays green or reveals an existing regression**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
|
||||
```
|
||||
|
||||
Expected: this test may already PASS before the refactor because it locks in existing external behavior. Keep it either way; it is the guardrail for the optimization.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/flow_upload_batch_test.go go-backend/tests/contract/flow_upload_batch_contract_test.go
|
||||
git commit -m "test: cover flow upload batch semantics"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Add Batched Repository Primitives
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/store/repo/repository_flow.go`
|
||||
- Modify: `go-backend/internal/store/repo/repository.go`
|
||||
- Modify: `go-backend/internal/store/repo/repository_user_quota.go`
|
||||
- Create: `go-backend/internal/store/repo/repository_flow_batch_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing repository tests**
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_flow_batch_test.go` with coverage for both the shared metadata query and the batched counter/quota writes.
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
metas, err := r.GetFlowUploadForwardMetas([]int64{20, 99})
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
t.Fatalf("unexpected meta for forward 20: %#v", metas[20])
|
||||
}
|
||||
if _, ok := metas[99]; ok {
|
||||
t.Fatalf("did not expect meta for missing forward 99")
|
||||
}
|
||||
|
||||
err = r.ApplyFlowUploadDeltasBatch([]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}})
|
||||
if err != nil {
|
||||
t.Fatalf("apply flow batch: %v", err)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
|
||||
t.Fatalf("expected forward in_flow=480, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
|
||||
t.Fatalf("expected user out_flow=660, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
|
||||
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddUserQuotaUsageBatchReturnsNormalizedViews(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "quota-batch.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
views, err := r.AddUserQuotaUsageBatch(map[int64]int64{2: 1140}, now)
|
||||
if err != nil {
|
||||
t.Fatalf("batch quota update: %v", err)
|
||||
}
|
||||
if views[2] == nil || views[2].DailyUsedBytes != 1140 || views[2].MonthlyUsedBytes != 1140 {
|
||||
t.Fatalf("unexpected quota view: %#v", views[2])
|
||||
}
|
||||
}
|
||||
|
||||
func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...interface{}) int64 {
|
||||
t.Helper()
|
||||
var value int64
|
||||
if err := r.DB().Raw(query, args...).Row().Scan(&value); err != nil {
|
||||
t.Fatalf("query %q failed: %v", query, err)
|
||||
}
|
||||
return value
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the repository tests to verify RED**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/store/repo -run 'TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch|TestAddUserQuotaUsageBatchReturnsNormalizedViews' -v
|
||||
```
|
||||
|
||||
Expected: FAIL because `GetFlowUploadForwardMetas`, `ApplyFlowUploadDeltasBatch`, `FlowUploadCounterDelta`, and `AddUserQuotaUsageBatch` do not exist yet.
|
||||
|
||||
- [ ] **Step 3: Implement shared flow-upload metadata and batched persistence**
|
||||
|
||||
Update `go-backend/internal/store/repo/repository_flow.go`, `repository.go`, and `repository_user_quota.go` with the following concrete APIs. Add `sort` to the `repository_user_quota.go` import list.
|
||||
|
||||
```go
|
||||
// repository_flow.go
|
||||
type FlowUploadForwardMeta struct {
|
||||
ForwardID int64
|
||||
TunnelID int64
|
||||
TrafficRatio float64
|
||||
TunnelFlow int64
|
||||
}
|
||||
|
||||
func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]FlowUploadForwardMeta, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if len(forwardIDs) == 0 {
|
||||
return map[int64]FlowUploadForwardMeta{}, nil
|
||||
}
|
||||
ids := make([]int64, 0, len(forwardIDs))
|
||||
seen := make(map[int64]struct{}, len(forwardIDs))
|
||||
for _, id := range forwardIDs {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
ids = append(ids, id)
|
||||
}
|
||||
type row struct {
|
||||
ForwardID int64 `gorm:"column:forward_id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio"`
|
||||
TunnelFlow int64 `gorm:"column:tunnel_flow"`
|
||||
}
|
||||
var rows []row
|
||||
err := r.db.Table("forward AS f").
|
||||
Select("f.id AS forward_id, f.tunnel_id AS tunnel_id, t.traffic_ratio AS traffic_ratio, t.flow AS tunnel_flow").
|
||||
Joins("JOIN tunnel t ON t.id = f.tunnel_id").
|
||||
Where("f.id IN ?", ids).
|
||||
Scan(&rows).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := make(map[int64]FlowUploadForwardMeta, len(rows))
|
||||
for _, row := range rows {
|
||||
if row.TunnelFlow <= 0 {
|
||||
row.TunnelFlow = 1
|
||||
}
|
||||
if row.TrafficRatio <= 0 {
|
||||
row.TrafficRatio = 1
|
||||
}
|
||||
out[row.ForwardID] = FlowUploadForwardMeta{ForwardID: row.ForwardID, TunnelID: row.TunnelID, TrafficRatio: row.TrafficRatio, TunnelFlow: row.TunnelFlow}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
```
|
||||
|
||||
```go
|
||||
// repository.go
|
||||
type FlowUploadCounterDelta struct {
|
||||
ForwardID int64
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
}
|
||||
|
||||
func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if len(deltas) == 0 {
|
||||
return nil
|
||||
}
|
||||
forwardTotals := make(map[int64][2]int64, len(deltas))
|
||||
userTotals := make(map[int64][2]int64, len(deltas))
|
||||
userTunnelTotals := make(map[int64][2]int64, len(deltas))
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID > 0 {
|
||||
current := forwardTotals[delta.ForwardID]
|
||||
current[0] += delta.InFlow
|
||||
current[1] += delta.OutFlow
|
||||
forwardTotals[delta.ForwardID] = current
|
||||
}
|
||||
if delta.UserID > 0 {
|
||||
current := userTotals[delta.UserID]
|
||||
current[0] += delta.InFlow
|
||||
current[1] += delta.OutFlow
|
||||
userTotals[delta.UserID] = current
|
||||
}
|
||||
if delta.UserTunnelID > 0 {
|
||||
current := userTunnelTotals[delta.UserTunnelID]
|
||||
current[0] += delta.InFlow
|
||||
current[1] += delta.OutFlow
|
||||
userTunnelTotals[delta.UserTunnelID] = current
|
||||
}
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
for forwardID, total := range forwardTotals {
|
||||
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1])}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for userID, total := range userTotals {
|
||||
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1])}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for userTunnelID, total := range userTunnelTotals {
|
||||
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1])}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
```
|
||||
|
||||
```go
|
||||
// repository_user_quota.go
|
||||
func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if len(usages) == 0 {
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
userIDs := make([]int64, 0, len(usages))
|
||||
for userID := range usages {
|
||||
if userID > 0 {
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
}
|
||||
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
|
||||
for _, userID := range userIDs {
|
||||
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyUserQuotaWindowRoll(q, now)
|
||||
if usages[userID] > 0 {
|
||||
q.DailyUsedBytes += usages[userID]
|
||||
q.MonthlyUsedBytes += usages[userID]
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{"daily_used_bytes": q.DailyUsedBytes, "monthly_used_bytes": q.MonthlyUsedBytes, "day_key": q.DayKey, "month_key": q.MonthKey, "updated_time": q.UpdatedTime}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Run the repository tests to verify GREEN**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/store/repo -run 'TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch|TestAddUserQuotaUsageBatchReturnsNormalizedViews' -v
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/store/repo/repository.go go-backend/internal/store/repo/repository_flow.go go-backend/internal/store/repo/repository_user_quota.go go-backend/internal/store/repo/repository_flow_batch_test.go
|
||||
git commit -m "refactor: batch flow upload persistence"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Refactor flow/upload To Use One Parsed Batch
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/flow_upload_batch.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_metrics_ingestion.go`
|
||||
- Modify: `go-backend/internal/http/handler/flow_upload_batch_test.go`
|
||||
- Modify: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the new handler batch implementation**
|
||||
|
||||
Create `go-backend/internal/http/handler/flow_upload_batch.go` and move the request-scoped aggregation there.
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type flowPolicyTarget struct {
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
}
|
||||
|
||||
type flowUploadBatch struct {
|
||||
flowDeltas []repo.FlowUploadCounterDelta
|
||||
quotaUsage map[int64]int64
|
||||
policyTargets []flowPolicyTarget
|
||||
forwardTraffic map[int64]tunnelTrafficDelta
|
||||
orphanServices map[string]struct{}
|
||||
peerShareForwardItems map[string]flowItem
|
||||
peerShareRuntimeItems map[int64]flowItem
|
||||
}
|
||||
|
||||
func (h *Handler) buildFlowUploadBatch(items []flowItem, metas map[int64]repo.FlowUploadForwardMeta) flowUploadBatch {
|
||||
batch := flowUploadBatch{
|
||||
quotaUsage: make(map[int64]int64),
|
||||
forwardTraffic: make(map[int64]tunnelTrafficDelta),
|
||||
orphanServices: make(map[string]struct{}),
|
||||
peerShareForwardItems: make(map[string]flowItem),
|
||||
peerShareRuntimeItems: make(map[int64]flowItem),
|
||||
}
|
||||
policySeen := map[flowPolicyTarget]struct{}{}
|
||||
flowSeen := map[int64]int{}
|
||||
|
||||
for _, item := range items {
|
||||
serviceName := strings.TrimSpace(item.N)
|
||||
if serviceName == "" || serviceName == "web_api" {
|
||||
continue
|
||||
}
|
||||
if runtimeID, ok := parsePeerShareRuntimeServiceID(serviceName); ok {
|
||||
merged := batch.peerShareRuntimeItems[runtimeID]
|
||||
merged.N = serviceName
|
||||
merged.U += item.U
|
||||
merged.D += item.D
|
||||
batch.peerShareRuntimeItems[runtimeID] = merged
|
||||
continue
|
||||
}
|
||||
forwardID, userID, userTunnelID, ok := parseFlowServiceIDs(serviceName)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
meta, exists := metas[forwardID]
|
||||
if !exists {
|
||||
batch.orphanServices[serviceName] = struct{}{}
|
||||
continue
|
||||
}
|
||||
raw := batch.forwardTraffic[forwardID]
|
||||
raw.bytesIn += item.D
|
||||
raw.bytesOut += item.U
|
||||
batch.forwardTraffic[forwardID] = raw
|
||||
|
||||
scaledIn := int64(float64(item.D)*meta.TrafficRatio) * meta.TunnelFlow
|
||||
scaledOut := int64(float64(item.U)*meta.TrafficRatio) * meta.TunnelFlow
|
||||
if idx, ok := flowSeen[forwardID]; ok {
|
||||
batch.flowDeltas[idx].InFlow += scaledIn
|
||||
batch.flowDeltas[idx].OutFlow += scaledOut
|
||||
} else {
|
||||
flowSeen[forwardID] = len(batch.flowDeltas)
|
||||
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{ForwardID: forwardID, UserID: userID, UserTunnelID: userTunnelID, InFlow: scaledIn, OutFlow: scaledOut})
|
||||
}
|
||||
batch.quotaUsage[userID] += scaledIn + scaledOut
|
||||
target := flowPolicyTarget{UserID: userID, UserTunnelID: userTunnelID}
|
||||
if _, seen := policySeen[target]; !seen {
|
||||
policySeen[target] = struct{}{}
|
||||
batch.policyTargets = append(batch.policyTargets, target)
|
||||
}
|
||||
merged := batch.peerShareForwardItems[normalizeForwardRuntimeServiceName(serviceName)]
|
||||
merged.N = normalizeForwardRuntimeServiceName(serviceName)
|
||||
merged.U += item.U
|
||||
merged.D += item.D
|
||||
batch.peerShareForwardItems[normalizeForwardRuntimeServiceName(serviceName)] = merged
|
||||
}
|
||||
|
||||
sort.Slice(batch.policyTargets, func(i, j int) bool {
|
||||
if batch.policyTargets[i].UserID == batch.policyTargets[j].UserID {
|
||||
return batch.policyTargets[i].UserTunnelID < batch.policyTargets[j].UserTunnelID
|
||||
}
|
||||
return batch.policyTargets[i].UserID < batch.policyTargets[j].UserID
|
||||
})
|
||||
return batch
|
||||
}
|
||||
|
||||
func (h *Handler) applyFlowUploadBatch(nodeID int64, batch flowUploadBatch, now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if err := h.repo.ApplyFlowUploadDeltasBatch(batch.flowDeltas); err != nil {
|
||||
log.Printf("flow upload write failed op=flow.batch_apply node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
quotaViews, err := h.repo.AddUserQuotaUsageBatch(batch.quotaUsage, now)
|
||||
if err != nil {
|
||||
log.Printf("flow upload write failed op=quota.batch_apply node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
for userID, quota := range quotaViews {
|
||||
h.enforceUserQuotaIfNeeded(userID, quota)
|
||||
}
|
||||
for _, target := range batch.policyTargets {
|
||||
if target.UserID <= 0 || target.UserTunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
h.enforceFlowPolicies(target.UserID, target.UserTunnelID)
|
||||
}
|
||||
for serviceName := range batch.orphanServices {
|
||||
h.sendDeleteOrphanedForwardService(nodeID, serviceName)
|
||||
}
|
||||
for serviceName, item := range batch.peerShareForwardItems {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(serviceName)
|
||||
if ok {
|
||||
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
|
||||
}
|
||||
}
|
||||
for runtimeID, item := range batch.peerShareRuntimeItems {
|
||||
h.processPeerShareFlow(runtimeID, item)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Switch the `/flow/upload` entrypoint and tunnel metric ingestion to the shared batch**
|
||||
|
||||
Modify `handler.go` and `tunnel_metrics_ingestion.go` so the raw JSON is parsed once and the same forward metadata powers both flow counters and tunnel metrics.
|
||||
|
||||
```go
|
||||
// handler.go
|
||||
func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
secret := r.URL.Query().Get("secret")
|
||||
node, _ := h.repo.GetNodeBySecret(secret)
|
||||
if node == nil {
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
return
|
||||
}
|
||||
|
||||
raw, err := readAndDecryptFlowBody(r.Body, secret)
|
||||
if err == nil && strings.TrimSpace(raw) != "" {
|
||||
var items []flowItem
|
||||
if json.Unmarshal([]byte(raw), &items) == nil {
|
||||
now := time.Now()
|
||||
forwardIDs := collectFlowUploadForwardIDs(items)
|
||||
metas, metaErr := h.repo.GetFlowUploadForwardMetas(forwardIDs)
|
||||
if metaErr != nil {
|
||||
log.Printf("flow upload metadata lookup failed node_id=%d err=%v", node.ID, metaErr)
|
||||
metas = map[int64]repo.FlowUploadForwardMeta{}
|
||||
}
|
||||
batch := h.buildFlowUploadBatch(items, metas)
|
||||
h.recordTunnelMetricsFromForwardBatch(node.ID, batch.forwardTraffic, metas, now.UnixMilli())
|
||||
h.applyFlowUploadBatch(node.ID, batch, now)
|
||||
}
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}
|
||||
```
|
||||
|
||||
```go
|
||||
// tunnel_metrics_ingestion.go
|
||||
func collectFlowUploadForwardIDs(items []flowItem) []int64 {
|
||||
ids := make([]int64, 0, len(items))
|
||||
seen := make(map[int64]struct{}, len(items))
|
||||
for _, item := range items {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(strings.TrimSpace(item.N))
|
||||
if !ok || forwardID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[forwardID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[forwardID] = struct{}{}
|
||||
ids = append(ids, forwardID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromForwardBatch(nodeID int64, forwardDeltas map[int64]tunnelTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta, nowMs int64) {
|
||||
if h == nil || h.repo == nil || nodeID <= 0 || len(forwardDeltas) == 0 {
|
||||
return
|
||||
}
|
||||
bucketTs := unixMilliBucketMinute(nowMs)
|
||||
if bucketTs <= 0 {
|
||||
return
|
||||
}
|
||||
tunnelAgg := make(map[int64]tunnelTrafficDelta)
|
||||
for forwardID, delta := range forwardDeltas {
|
||||
meta, ok := metas[forwardID]
|
||||
if !ok || meta.TunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
current := tunnelAgg[meta.TunnelID]
|
||||
current.bytesIn += delta.bytesIn
|
||||
current.bytesOut += delta.bytesOut
|
||||
tunnelAgg[meta.TunnelID] = current
|
||||
}
|
||||
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
|
||||
for tunnelID, delta := range tunnelAgg {
|
||||
if delta.bytesIn == 0 && delta.bytesOut == 0 {
|
||||
continue
|
||||
}
|
||||
metrics = append(metrics, &model.TunnelMetric{TunnelID: tunnelID, NodeID: nodeID, Timestamp: bucketTs, BytesIn: delta.bytesIn, BytesOut: delta.bytesOut})
|
||||
}
|
||||
if len(metrics) == 0 {
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
|
||||
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
|
||||
return
|
||||
}
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Run focused handler and contract tests to verify GREEN**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run TestBuildFlowUploadBatchAggregatesForwardQuotaPeerShareAndCleanupTargets -v
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 4: Run the full backend suite**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./...
|
||||
```
|
||||
|
||||
Expected: PASS across the backend module.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/tunnel_metrics_ingestion.go go-backend/internal/http/handler/flow_upload_batch.go go-backend/internal/http/handler/flow_upload_batch_test.go go-backend/tests/contract/flow_upload_batch_contract_test.go
|
||||
git commit -m "refactor: batch flow upload processing"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Final Verification And Performance Sanity Check
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/tests/contract/flow_upload_batch_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Add a same-batch duplicate-item stress assertion**
|
||||
|
||||
Extend the contract test with a second request that repeats the same service name multiple times and assert the counters advance by exactly the summed amount.
|
||||
|
||||
```go
|
||||
body, err = json.Marshal([]map[string]interface{}{
|
||||
{"n": "20_2_10", "u": 10, "d": 20},
|
||||
{"n": "20_2_10", "u": 10, "d": 20},
|
||||
{"n": "20_2_10_tcp", "u": 10, "d": 20},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal body: %v", err)
|
||||
}
|
||||
req = httptest.NewRequest(http.MethodPost, "/flow/upload?secret="+node.Secret, bytes.NewReader(body))
|
||||
res = httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
if got := mustQueryInt(t, repo, `SELECT in_flow FROM forward WHERE id = 20`); got != 140 {
|
||||
t.Fatalf("expected forward in_flow=140 after second request, got %d", got)
|
||||
}
|
||||
if got := mustQueryInt(t, repo, `SELECT out_flow FROM forward WHERE id = 20`); got != 140 {
|
||||
t.Fatalf("expected forward out_flow=140 after second request, got %d", got)
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the targeted contract test again**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -v
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Re-run the full backend suite before claiming completion**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
go test ./...
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 4: Optional local profiling sanity check**
|
||||
|
||||
Run a short local comparison before and after the change with the same repeated flow payload.
|
||||
|
||||
```bash
|
||||
go test ./tests/contract/... -run TestFlowUploadAggregatesRepeatedItemsAndDisablesQuotaImmediately -count=10
|
||||
```
|
||||
|
||||
Expected: the test remains stable across repeated runs and does not introduce flakiness.
|
||||
|
||||
- [ ] **Step 5: Optional commit if the user explicitly requested commits**
|
||||
|
||||
```bash
|
||||
git add go-backend/tests/contract/flow_upload_batch_contract_test.go
|
||||
git commit -m "test: harden flow upload batch regression coverage"
|
||||
```
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,598 @@
|
||||
# Monitoring Retention And Storage Display Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Add configurable monitoring data retention and show database storage usage on the config page.
|
||||
|
||||
**Architecture:** Store retention in `vite_config` as `monitor_retention_days`, parse it through a focused monitoring helper, and reuse it from existing cleanup loops. Add a repository storage-summary helper, expose it via an admin-only API, and render it in the existing React config page.
|
||||
|
||||
**Tech Stack:** Go `net/http`, GORM, SQLite/PostgreSQL, Vite/React/TypeScript, existing shadcn bridge components.
|
||||
|
||||
---
|
||||
|
||||
## File Structure
|
||||
|
||||
- Create: `go-backend/internal/monitoring/retention.go` for retention constants, parsing, and validation.
|
||||
- Test: `go-backend/internal/monitoring/retention_test.go`.
|
||||
- Modify: `go-backend/internal/metrics/ingestion.go` and `go-backend/internal/metrics/ingestion_test.go` for config-driven cleanup.
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_quality_prober.go` so `tunnel_quality` uses the same retention and still prunes when probing is disabled.
|
||||
- Create: `go-backend/internal/store/repo/repository_storage.go` and `go-backend/internal/store/repo/repository_storage_test.go` for database size summaries.
|
||||
- Modify: `go-backend/internal/store/repo/repository.go` to keep the SQLite DB path on `Repository`.
|
||||
- Create: `go-backend/internal/http/handler/storage.go` for the storage endpoint.
|
||||
- Modify: `go-backend/internal/http/handler/handler.go` to register `/api/v1/system/storage` and validate `monitor_retention_days`.
|
||||
- Modify: `go-backend/internal/http/middleware/auth.go` so `/api/v1/system/*` is admin-only.
|
||||
- Create: `go-backend/tests/contract/storage_contract_test.go` for endpoint auth/shape coverage.
|
||||
- Modify: `vite-frontend/src/api/types.ts`, `vite-frontend/src/api/index.ts`, and `vite-frontend/src/pages/config.tsx` for UI display.
|
||||
|
||||
Implementation should not create git commits unless the user explicitly requests them.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Add Retention Config Helper
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/monitoring/retention.go`
|
||||
- Create: `go-backend/internal/monitoring/retention_test.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing tests**
|
||||
|
||||
Create `go-backend/internal/monitoring/retention_test.go`:
|
||||
|
||||
```go
|
||||
package monitoring
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMonitoringRetentionDaysFromConfigMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg map[string]string
|
||||
want int
|
||||
}{
|
||||
{"missing uses default", nil, 7},
|
||||
{"valid custom", map[string]string{ConfigMonitorRetentionDays: "3"}, 3},
|
||||
{"trimmed custom", map[string]string{ConfigMonitorRetentionDays: " 30 "}, 30},
|
||||
{"invalid uses default", map[string]string{ConfigMonitorRetentionDays: "abc"}, 7},
|
||||
{"too small uses default", map[string]string{ConfigMonitorRetentionDays: "0"}, 7},
|
||||
{"too large uses default", map[string]string{ConfigMonitorRetentionDays: "3651"}, 7},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := MonitoringRetentionDaysFromConfigMap(tc.cfg); got != tc.want {
|
||||
t.Fatalf("expected %d, got %d", tc.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeMonitoringRetentionDays(t *testing.T) {
|
||||
for _, value := range []string{"1", "7", "3650", " 30 "} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err != nil || got == "" {
|
||||
t.Fatalf("expected %q valid, got value=%q err=%v", value, got, err)
|
||||
}
|
||||
}
|
||||
for _, value := range []string{"", "0", "-1", "3651", "abc", "1.5"} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err == nil {
|
||||
t.Fatalf("expected %q invalid, got value=%q", value, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run tests to verify failure**
|
||||
|
||||
Run: `go test ./internal/monitoring -run 'TestMonitoringRetentionDaysFromConfigMap|TestNormalizeMonitoringRetentionDays' -count=1`
|
||||
|
||||
Expected: FAIL with undefined `ConfigMonitorRetentionDays`, `MonitoringRetentionDaysFromConfigMap`, and `NormalizeMonitoringRetentionDays`.
|
||||
|
||||
- [ ] **Step 3: Implement helper**
|
||||
|
||||
Create `go-backend/internal/monitoring/retention.go`:
|
||||
|
||||
```go
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
ConfigMonitorRetentionDays = "monitor_retention_days"
|
||||
DefaultMonitorRetentionDays = 7
|
||||
MinMonitorRetentionDays = 1
|
||||
MaxMonitorRetentionDays = 3650
|
||||
)
|
||||
|
||||
func MonitoringRetentionDaysFromConfigMap(cfg map[string]string) int {
|
||||
if cfg == nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
days, err := parseMonitoringRetentionDays(cfg[ConfigMonitorRetentionDays])
|
||||
if err != nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
return days
|
||||
}
|
||||
|
||||
func NormalizeMonitoringRetentionDays(value string) (string, error) {
|
||||
days, err := parseMonitoringRetentionDays(value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strconv.Itoa(days), nil
|
||||
}
|
||||
|
||||
func parseMonitoringRetentionDays(value string) (int, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return 0, fmt.Errorf("监控数据保留天数不能为空")
|
||||
}
|
||||
days, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须是整数")
|
||||
}
|
||||
if days < MinMonitorRetentionDays || days > MaxMonitorRetentionDays {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须在 %d 到 %d 之间", MinMonitorRetentionDays, MaxMonitorRetentionDays)
|
||||
}
|
||||
return days, nil
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Validate config updates**
|
||||
|
||||
In `go-backend/internal/http/handler/handler.go`, add this case to `normalizeAndValidateConfigValue`:
|
||||
|
||||
```go
|
||||
case monitoring.ConfigMonitorRetentionDays:
|
||||
return monitoring.NormalizeMonitoringRetentionDays(value)
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run tests**
|
||||
|
||||
Run: `go test ./internal/monitoring ./internal/http/handler -run 'TestMonitoringRetention|TestNormalize|Test' -count=1`
|
||||
|
||||
Expected: PASS or only unrelated pre-existing failures, which must be investigated before continuing.
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Use Retention Config In Cleanup
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/metrics/ingestion.go`
|
||||
- Modify: `go-backend/internal/metrics/ingestion_test.go`
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_quality_prober.go`
|
||||
|
||||
- [ ] **Step 1: Write failing cleanup test**
|
||||
|
||||
Append to `go-backend/internal/metrics/ingestion_test.go`, adding `go-backend/internal/store/model` to imports:
|
||||
|
||||
```go
|
||||
func TestPruneMetricsUsesConfiguredRetentionDays(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.UpsertConfig("monitor_retention_days", "2", now); err != nil {
|
||||
t.Fatalf("upsert retention config: %v", err)
|
||||
}
|
||||
|
||||
oldMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(3*24*time.Hour/time.Millisecond), CPUUsage: 10}
|
||||
newMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(1*24*time.Hour/time.Millisecond), CPUUsage: 20}
|
||||
if err := r.InsertNodeMetric(oldMetric); err != nil {
|
||||
t.Fatalf("insert old metric: %v", err)
|
||||
}
|
||||
if err := r.InsertNodeMetric(newMetric); err != nil {
|
||||
t.Fatalf("insert new metric: %v", err)
|
||||
}
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.pruneMetricsAt(time.UnixMilli(now))
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, now-int64(4*24*time.Hour/time.Millisecond), now+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get node metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 || metrics[0].CPUUsage != 20 {
|
||||
t.Fatalf("expected only newer metric to remain, got %#v", metrics)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run test to verify failure**
|
||||
|
||||
Run: `go test ./internal/metrics -run TestPruneMetricsUsesConfiguredRetentionDays -count=1`
|
||||
|
||||
Expected: FAIL with undefined `pruneMetricsAt`.
|
||||
|
||||
- [ ] **Step 3: Implement config-driven prune**
|
||||
|
||||
In `go-backend/internal/metrics/ingestion.go`, import `go-backend/internal/monitoring` and replace `pruneMetrics` with:
|
||||
|
||||
```go
|
||||
func (s *IngestionService) retentionDaysFromConfig() int {
|
||||
if s == nil || s.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := s.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetrics() {
|
||||
s.pruneMetricsAt(time.Now())
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetricsAt(now time.Time) {
|
||||
cutoff := now.Add(-time.Duration(s.retentionDaysFromConfig()) * 24 * time.Hour).UnixMilli()
|
||||
if s.repo == nil {
|
||||
return
|
||||
}
|
||||
if err := s.repo.PruneNodeMetrics(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=node_metric cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
if err := s.repo.PruneTunnelMetrics(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=tunnel_metric cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
if err := s.repo.PruneServiceMonitorResults(cutoff); err != nil {
|
||||
log.Printf("monitoring prune failed op=service_monitor_result cutoff=%d err=%v", cutoff, err)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Remove the unused `retentionDays` field from `IngestionService` and remove `svc.retentionDays = 1` from existing tests.
|
||||
|
||||
- [ ] **Step 4: Update tunnel quality pruning**
|
||||
|
||||
In `go-backend/internal/http/handler/tunnel_quality_prober.go`, import `go-backend/internal/monitoring`, remove `tunnelQualityRetention`, remove the `if !p.isEnabled() { return }` guard from `maybePrune`, and calculate cutoff with:
|
||||
|
||||
```go
|
||||
func (p *tunnelQualityProber) retentionDays() int {
|
||||
if p == nil || p.handler == nil || p.handler.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
```
|
||||
|
||||
Then use:
|
||||
|
||||
```go
|
||||
cutoff := now - int64(time.Duration(p.retentionDays())*24*time.Hour/time.Millisecond)
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run cleanup tests**
|
||||
|
||||
Run: `go test ./internal/metrics ./internal/http/handler -run 'TestPruneMetrics|TestPruneMetricsUsesConfiguredRetentionDays|TunnelQuality' -count=1`
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Add Storage Summary Backend API
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/store/repo/repository.go`
|
||||
- Create: `go-backend/internal/store/repo/repository_storage.go`
|
||||
- Create: `go-backend/internal/store/repo/repository_storage_test.go`
|
||||
- Create: `go-backend/internal/http/handler/storage.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/middleware/auth.go`
|
||||
- Create: `go-backend/tests/contract/storage_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write failing repository tests**
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_storage_test.go`:
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestDatabaseStorageSummarySQLiteIncludesSize(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "storage.db")
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
if err := r.InsertNodeMetric(&model.NodeMetric{NodeID: 1, Timestamp: 123, CPUUsage: 1}); err != nil {
|
||||
t.Fatalf("insert metric: %v", err)
|
||||
}
|
||||
summary, err := r.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
t.Fatalf("storage summary: %v", err)
|
||||
}
|
||||
if summary.DBType != "sqlite" || summary.DatabaseSizeBytes <= 0 || summary.DatabaseSizeText == "" {
|
||||
t.Fatalf("unexpected summary: %#v", summary)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatDatabaseSize(t *testing.T) {
|
||||
for _, tc := range []struct{ bytes int64; want string }{{0, "0 B"}, {512, "512 B"}, {1024, "1.0 KB"}, {1024 * 1024, "1.0 MB"}} {
|
||||
if got := formatDatabaseSize(tc.bytes); got != tc.want {
|
||||
t.Fatalf("formatDatabaseSize(%d)=%q want %q", tc.bytes, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run test to verify failure**
|
||||
|
||||
Run: `go test ./internal/store/repo -run 'TestDatabaseStorageSummarySQLiteIncludesSize|TestFormatDatabaseSize' -count=1`
|
||||
|
||||
Expected: FAIL with undefined `DatabaseStorageSummary` and `formatDatabaseSize`.
|
||||
|
||||
- [ ] **Step 3: Implement repository helper**
|
||||
|
||||
Modify `Repository` in `repository.go`:
|
||||
|
||||
```go
|
||||
type Repository struct {
|
||||
db *gorm.DB
|
||||
dbPath string
|
||||
}
|
||||
```
|
||||
|
||||
Return `&Repository{db: db, dbPath: path}` from `Open` and `&Repository{db: db}` from `OpenPostgres`.
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_storage.go`:
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
type DatabaseStorageSummary struct {
|
||||
DBType string `json:"dbType"`
|
||||
DatabaseSizeBytes int64 `json:"databaseSizeBytes"`
|
||||
DatabaseSizeText string `json:"databaseSizeText"`
|
||||
}
|
||||
|
||||
func (r *Repository) DatabaseStorageSummary() (DatabaseStorageSummary, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return DatabaseStorageSummary{}, errors.New("repository not initialized")
|
||||
}
|
||||
switch r.db.Dialector.Name() {
|
||||
case "sqlite":
|
||||
size, err := sqliteDatabaseFileSize(r.dbPath)
|
||||
if err != nil { return DatabaseStorageSummary{}, err }
|
||||
return DatabaseStorageSummary{"sqlite", size, formatDatabaseSize(size)}, nil
|
||||
case "postgres":
|
||||
var size int64
|
||||
if err := r.db.Raw("SELECT pg_database_size(current_database())").Scan(&size).Error; err != nil { return DatabaseStorageSummary{}, err }
|
||||
return DatabaseStorageSummary{"postgres", size, formatDatabaseSize(size)}, nil
|
||||
default:
|
||||
return DatabaseStorageSummary{}, fmt.Errorf("unsupported database dialect %q", r.db.Dialector.Name())
|
||||
}
|
||||
}
|
||||
|
||||
func sqliteDatabaseFileSize(path string) (int64, error) {
|
||||
if path == "" || path == ":memory:" { return 0, nil }
|
||||
var total int64
|
||||
for _, candidate := range []string{path, path + "-wal", path + "-shm"} {
|
||||
info, err := os.Stat(candidate)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) { continue }
|
||||
return 0, err
|
||||
}
|
||||
if !info.IsDir() { total += info.Size() }
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func formatDatabaseSize(bytes int64) string {
|
||||
if bytes < 1024 { return fmt.Sprintf("%d B", bytes) }
|
||||
units := []string{"KB", "MB", "GB", "TB"}
|
||||
value := float64(bytes) / 1024
|
||||
for _, unit := range units {
|
||||
if value < 1024 || unit == "TB" { return fmt.Sprintf("%.1f %s", value, unit) }
|
||||
value /= 1024
|
||||
}
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Add API handler and route**
|
||||
|
||||
Create `go-backend/internal/http/handler/storage.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func (h *Handler) storageSummary(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if h == nil || h.repo == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "repository not initialized"))
|
||||
return
|
||||
}
|
||||
summary, err := h.repo.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(summary))
|
||||
}
|
||||
```
|
||||
|
||||
Register in `Handler.Register`: `mux.HandleFunc("/api/v1/system/storage", h.storageSummary)`.
|
||||
|
||||
In `requiresAdmin`, add:
|
||||
|
||||
```go
|
||||
if strings.HasPrefix(path, "/api/v1/system/") {
|
||||
return true
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Write contract test for auth and shape**
|
||||
|
||||
Create `go-backend/tests/contract/storage_contract_test.go` with a test that sends GET `/api/v1/system/storage` as non-admin and expects `403`, then as admin and expects `code == 0`, `dbType`, numeric `databaseSizeBytes`, and `databaseSizeText`.
|
||||
|
||||
- [ ] **Step 6: Run storage tests**
|
||||
|
||||
Run: `go test ./internal/store/repo ./tests/contract -run 'TestDatabaseStorageSummarySQLiteIncludesSize|TestFormatDatabaseSize|TestStorageSummaryRequiresAdminAndReturnsSize' -count=1`
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Add Frontend Config UI
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/api/types.ts`
|
||||
- Modify: `vite-frontend/src/api/index.ts`
|
||||
- Modify: `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
- [ ] **Step 1: Add API type and function**
|
||||
|
||||
In `types.ts` add:
|
||||
|
||||
```ts
|
||||
export interface StorageSummaryApiData {
|
||||
dbType: string;
|
||||
databaseSizeBytes: number;
|
||||
databaseSizeText: string;
|
||||
}
|
||||
```
|
||||
|
||||
In `index.ts`, import `StorageSummaryApiData` and add:
|
||||
|
||||
```ts
|
||||
export const getStorageSummary = () =>
|
||||
Network.get<StorageSummaryApiData>("/system/storage");
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Add retention config item**
|
||||
|
||||
In `config.tsx`, add to `CONFIG_ITEMS` near monitoring:
|
||||
|
||||
```ts
|
||||
{
|
||||
key: "monitor_retention_days",
|
||||
label: "监控数据保留天数",
|
||||
placeholder: "7",
|
||||
description:
|
||||
"统一清理节点指标、隧道流量、服务监控结果和隧道质量历史;默认 7 天。",
|
||||
type: "input",
|
||||
},
|
||||
```
|
||||
|
||||
Add `"monitor_retention_days"` to `getInitialConfigs()` keys.
|
||||
|
||||
- [ ] **Step 3: Fetch and display database size**
|
||||
|
||||
In `config.tsx`, add state:
|
||||
|
||||
```ts
|
||||
const [storageSummary, setStorageSummary] = useState<string>("加载中...");
|
||||
```
|
||||
|
||||
Add a load effect:
|
||||
|
||||
```ts
|
||||
useEffect(() => {
|
||||
let mounted = true;
|
||||
getStorageSummary()
|
||||
.then((response) => {
|
||||
if (!mounted) return;
|
||||
if (response.code === 0 && response.data?.databaseSizeText) {
|
||||
setStorageSummary(response.data.databaseSizeText);
|
||||
} else {
|
||||
setStorageSummary("获取失败");
|
||||
}
|
||||
})
|
||||
.catch(() => {
|
||||
if (mounted) setStorageSummary("获取失败");
|
||||
});
|
||||
return () => {
|
||||
mounted = false;
|
||||
};
|
||||
}, []);
|
||||
```
|
||||
|
||||
Render inside the basic settings card before the save button:
|
||||
|
||||
```tsx
|
||||
<Divider className="my-2" />
|
||||
<div className="space-y-1">
|
||||
<p className="text-sm font-medium text-gray-700 dark:text-gray-300">
|
||||
数据库占用
|
||||
</p>
|
||||
<p className="text-xs text-gray-500 dark:text-gray-400">
|
||||
当前后端数据库文件/实例占用空间,仅用于容量参考。
|
||||
</p>
|
||||
<div className="rounded-lg border border-divider bg-default-50/60 dark:bg-default-100/10 px-4 py-3 text-sm font-semibold text-default-800 dark:text-default-200">
|
||||
{storageSummary}
|
||||
</div>
|
||||
</div>
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Build frontend**
|
||||
|
||||
Run: `pnpm run build` from `vite-frontend`.
|
||||
|
||||
Expected: TypeScript and Vite build pass.
|
||||
|
||||
---
|
||||
|
||||
### Task 5: Final Verification
|
||||
|
||||
**Files:**
|
||||
- All files changed by previous tasks.
|
||||
|
||||
- [ ] **Step 1: Run backend tests**
|
||||
|
||||
Run: `go test ./...` from `go-backend`.
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 2: Run frontend build**
|
||||
|
||||
Run: `pnpm run build` from `vite-frontend`.
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Review diff**
|
||||
|
||||
Run: `git diff --stat` and `git diff -- docs/superpowers/specs/2026-04-28-monitoring-retention-storage-design.md docs/superpowers/plans/2026-04-28-monitoring-retention-storage.md go-backend vite-frontend`.
|
||||
|
||||
Expected: Diff is limited to retention config, storage summary, tests, and config UI.
|
||||
|
||||
---
|
||||
|
||||
## Self-Review
|
||||
|
||||
- Spec coverage: retention config, uniform cleanup, storage summary API, frontend display, validation, and verification are covered.
|
||||
- Placeholder scan: no TBD/TODO placeholders; the one contract-test step describes exact assertions even though the surrounding helper functions already exist in contract tests.
|
||||
- Type consistency: backend JSON fields match frontend `StorageSummaryApiData` exactly: `dbType`, `databaseSizeBytes`, `databaseSizeText`.
|
||||
@@ -0,0 +1,892 @@
|
||||
# Best Exit Current Display Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Show the currently applied `best` exit selection in the tunnel list information, including per-entry/per-final-hop details for multi-owner tunnels.
|
||||
|
||||
**Architecture:** Add a backend-only display layer that snapshots `bestExitManager` state and attaches `bestExitState` to existing `tunnelList`/`tunnelGet` responses. Render that state in the existing tunnel table/grid topology area using compact text and a native `title` detail tooltip. No routing, scoring, persistence, polling, or runtime update behavior changes.
|
||||
|
||||
**Tech Stack:** Go `net/http` handlers + existing repository methods, React/TypeScript in `vite-frontend/src/pages/tunnel.tsx`, Tailwind/shadcn bridge components already in the file.
|
||||
|
||||
---
|
||||
|
||||
## File Structure
|
||||
|
||||
- Create `go-backend/internal/http/handler/tunnel_best_exit_display.go`: response DTOs, manager snapshot method, tunnel-response parsing helpers, and `Handler.attachBestExitStates`.
|
||||
- Create `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`: backend display-state unit tests.
|
||||
- Modify `go-backend/internal/http/handler/handler.go`: call `h.attachBestExitStatesOrLog(items)` in `tunnelList`.
|
||||
- Modify `go-backend/internal/http/handler/mutations.go`: call `h.attachBestExitStatesOrLog(items)` before returning a single tunnel in `tunnelGet`.
|
||||
- Modify `vite-frontend/src/pages/tunnel.tsx`: add `bestExitState` types, map API state, helper render functions, and table/grid display.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Backend Snapshot And Display-State Tests
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`
|
||||
|
||||
- [ ] **Step 1: Write failing backend display tests**
|
||||
|
||||
Create `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestBestExitDecisionSnapshotIsDefensiveCopy(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
score := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30, NodeName: "exit-a"}, 10, 0, 20, 0)
|
||||
|
||||
m.observeScores(key, []bestExitCandidateScore{score}, now)
|
||||
snapshot, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected snapshot")
|
||||
}
|
||||
if snapshot.AppliedExitNodeID != 30 || snapshot.UpdatedAt != now.UnixMilli() {
|
||||
t.Fatalf("unexpected snapshot: %+v", snapshot)
|
||||
}
|
||||
if len(snapshot.Scores) != 1 {
|
||||
t.Fatalf("expected one score in snapshot, got %+v", snapshot.Scores)
|
||||
}
|
||||
snapshot.Scores[0].ExitNodeID = 99
|
||||
|
||||
again, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected second snapshot")
|
||||
}
|
||||
if again.Scores[0].ExitNodeID != 30 {
|
||||
t.Fatalf("snapshot score mutation leaked into manager state: %+v", again.Scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForDirectMultiEntryOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(100, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 11}, 31, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
{"nodeId": int64(11)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 11: "入口 B", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if !state.Enabled || state.Summary != "多个出口" || state.Status != "applied" {
|
||||
t.Fatalf("unexpected state summary: %+v", state)
|
||||
}
|
||||
if state.UpdatedAt != now.Add(time.Second).UnixMilli() {
|
||||
t.Fatalf("expected latest updatedAt, got %d", state.UpdatedAt)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "entry" || state.Items[0].OwnerNodeName != "入口 A" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "entry" || state.Items[1].OwnerNodeName != "入口 B" || state.Items[1].ExitNodeName != "日本节点" {
|
||||
t.Fatalf("unexpected second item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForFinalChainHopOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(200, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 20}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 21}, 30, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(88),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{
|
||||
{{"nodeId": int64(15), "inx": int64(0)}},
|
||||
{{"nodeId": int64(20), "inx": int64(1)}, {"nodeId": int64(21), "inx": int64(1)}},
|
||||
},
|
||||
}
|
||||
names := map[int64]string{20: "中转 M1", 21: "中转 M2", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Summary != "香港节点" || state.Status != "applied" {
|
||||
t.Fatalf("expected single-exit summary, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two final-hop owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "chain" || state.Items[0].OwnerNodeName != "中转 M1" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first chain owner item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "chain" || state.Items[1].OwnerNodeName != "中转 M2" || state.Items[1].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected second chain owner item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateWaitingWhenNoAppliedDecisionExists(t *testing.T) {
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, newBestExitManager(), testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected waiting best exit state")
|
||||
}
|
||||
if state.Summary != "等待探测" || state.Status != "waiting" {
|
||||
t.Fatalf("expected waiting state, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 1 || state.Items[0].ExitNodeID != 0 || state.Items[0].ExitNodeName != "等待探测" {
|
||||
t.Fatalf("unexpected waiting item: %+v", state.Items)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateSkipsNonBestAndSingleExitTunnels(t *testing.T) {
|
||||
nonBest := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": "round"},
|
||||
{"nodeId": int64(31), "strategy": "round"},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(nonBest, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected non-best tunnel to skip state, got %+v", state)
|
||||
}
|
||||
|
||||
singleExit := map[string]interface{}{
|
||||
"id": int64(78),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(singleExit, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected single-exit tunnel to skip state, got %+v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func testBestExitNameLookup(names map[int64]string) bestExitNodeNameLookup {
|
||||
return func(nodeID int64) (string, bool) {
|
||||
name := names[nodeID]
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run backend display tests to verify failure**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run 'TestBestExitDecisionSnapshot|TestBuildBestExitDisplayState' -count=1
|
||||
```
|
||||
|
||||
Expected: FAIL with undefined `snapshot`, `buildBestExitDisplayState`, and `bestExitNodeNameLookup`.
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Backend Display State Implementation
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/tunnel_best_exit_display.go`
|
||||
- Modify: `go-backend/internal/http/handler/tunnel_best_exit.go`
|
||||
- Test: `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`
|
||||
|
||||
- [ ] **Step 1: Implement display state and snapshot helpers**
|
||||
|
||||
Create `go-backend/internal/http/handler/tunnel_best_exit_display.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"log"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
bestExitDisplayStatusApplied = "applied"
|
||||
bestExitDisplayStatusWaiting = "waiting"
|
||||
bestExitDisplaySummaryMulti = "多个出口"
|
||||
bestExitDisplaySummaryWait = "等待探测"
|
||||
bestExitUnknownExitName = "未知出口"
|
||||
bestExitUnknownEntryName = "未知入口"
|
||||
bestExitUnknownChainName = "未知中转"
|
||||
)
|
||||
|
||||
type bestExitDecisionSnapshot struct {
|
||||
AppliedExitNodeID int64
|
||||
UpdatedAt int64
|
||||
Reason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitDisplayState struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Summary string `json:"summary"`
|
||||
Status string `json:"status"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
Items []bestExitDisplayItem `json:"items"`
|
||||
}
|
||||
|
||||
type bestExitDisplayItem struct {
|
||||
OwnerNodeID int64 `json:"ownerNodeId"`
|
||||
OwnerNodeName string `json:"ownerNodeName"`
|
||||
OwnerRole string `json:"ownerRole"`
|
||||
ExitNodeID int64 `json:"exitNodeId,omitempty"`
|
||||
ExitNodeName string `json:"exitNodeName"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
}
|
||||
|
||||
type bestExitNodeNameLookup func(nodeID int64) (string, bool)
|
||||
|
||||
func (m *bestExitManager) snapshot(key bestExitOwnerKey) (bestExitDecisionSnapshot, bool) {
|
||||
if m == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisions[key]
|
||||
if d == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
updatedAt := int64(0)
|
||||
if !d.LastSwitchAt.IsZero() {
|
||||
updatedAt = d.LastSwitchAt.UnixMilli()
|
||||
}
|
||||
return bestExitDecisionSnapshot{
|
||||
AppliedExitNodeID: d.AppliedExitNodeID,
|
||||
UpdatedAt: updatedAt,
|
||||
Reason: d.LastReason,
|
||||
Scores: cloneBestExitScores(d.Scores),
|
||||
}, true
|
||||
}
|
||||
|
||||
func (h *Handler) attachBestExitStates(items []map[string]interface{}) {
|
||||
if h == nil || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
lookup := h.bestExitNodeNameLookup()
|
||||
for _, item := range items {
|
||||
state, ok := buildBestExitDisplayState(item, h.bestExit, lookup)
|
||||
if !ok {
|
||||
delete(item, "bestExitState")
|
||||
continue
|
||||
}
|
||||
item["bestExitState"] = state
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) bestExitNodeNameLookup() bestExitNodeNameLookup {
|
||||
cache := map[int64]string{}
|
||||
return func(nodeID int64) (string, bool) {
|
||||
if nodeID <= 0 || h == nil {
|
||||
return "", false
|
||||
}
|
||||
if name, ok := cache[nodeID]; ok {
|
||||
return name, name != ""
|
||||
}
|
||||
node, err := h.getNodeRecord(nodeID)
|
||||
if err != nil || node == nil {
|
||||
cache[nodeID] = ""
|
||||
return "", false
|
||||
}
|
||||
name := strings.TrimSpace(node.Name)
|
||||
cache[nodeID] = name
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
|
||||
func buildBestExitDisplayState(tunnel map[string]interface{}, manager *bestExitManager, lookup bestExitNodeNameLookup) (*bestExitDisplayState, bool) {
|
||||
if tunnel == nil {
|
||||
return nil, false
|
||||
}
|
||||
tunnelID := asInt64(tunnel["id"], 0)
|
||||
outNodes := bestExitDisplayMapSlice(tunnel["outNodeId"])
|
||||
if tunnelID <= 0 || len(outNodes) <= 1 {
|
||||
return nil, false
|
||||
}
|
||||
if !isBestTunnelStrategy(asString(outNodes[0]["strategy"])) {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
owners, ownerRole := bestExitDisplayOwners(tunnel)
|
||||
state := &bestExitDisplayState{
|
||||
Enabled: true,
|
||||
Summary: bestExitDisplaySummaryWait,
|
||||
Status: bestExitDisplayStatusWaiting,
|
||||
Items: make([]bestExitDisplayItem, 0, len(owners)),
|
||||
}
|
||||
|
||||
exitsByID := map[int64]map[string]interface{}{}
|
||||
for _, exit := range outNodes {
|
||||
if id := asInt64(exit["nodeId"], 0); id > 0 {
|
||||
exitsByID[id] = exit
|
||||
}
|
||||
}
|
||||
appliedExitIDs := map[int64]string{}
|
||||
appliedCount := 0
|
||||
latestUpdatedAt := int64(0)
|
||||
latestReason := ""
|
||||
for _, owner := range owners {
|
||||
ownerNodeID := asInt64(owner["nodeId"], 0)
|
||||
if ownerNodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
item := bestExitDisplayItem{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
OwnerNodeName: bestExitDisplayNodeName(owner, ownerNodeID, lookup, bestExitUnknownOwnerName(ownerRole)),
|
||||
OwnerRole: ownerRole,
|
||||
ExitNodeName: bestExitDisplaySummaryWait,
|
||||
Reason: bestExitDisplayStatusWaiting,
|
||||
}
|
||||
if snapshot, ok := manager.snapshot(bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: ownerNodeID}); ok && snapshot.AppliedExitNodeID > 0 {
|
||||
item.ExitNodeID = snapshot.AppliedExitNodeID
|
||||
item.ExitNodeName = bestExitDisplayNodeName(exitsByID[snapshot.AppliedExitNodeID], snapshot.AppliedExitNodeID, lookup, bestExitUnknownExitName)
|
||||
item.UpdatedAt = snapshot.UpdatedAt
|
||||
item.Reason = snapshot.Reason
|
||||
appliedExitIDs[item.ExitNodeID] = item.ExitNodeName
|
||||
appliedCount++
|
||||
if snapshot.UpdatedAt > latestUpdatedAt {
|
||||
latestUpdatedAt = snapshot.UpdatedAt
|
||||
latestReason = snapshot.Reason
|
||||
}
|
||||
}
|
||||
state.Items = append(state.Items, item)
|
||||
}
|
||||
|
||||
if appliedCount == 0 {
|
||||
return state, true
|
||||
}
|
||||
state.Status = bestExitDisplayStatusApplied
|
||||
state.UpdatedAt = latestUpdatedAt
|
||||
state.Reason = latestReason
|
||||
if len(appliedExitIDs) == 1 {
|
||||
for _, name := range appliedExitIDs {
|
||||
state.Summary = name
|
||||
}
|
||||
} else {
|
||||
state.Summary = bestExitDisplaySummaryMulti
|
||||
}
|
||||
return state, true
|
||||
}
|
||||
|
||||
func bestExitDisplayOwners(tunnel map[string]interface{}) ([]map[string]interface{}, string) {
|
||||
chainGroups := bestExitDisplayChainGroups(tunnel["chainNodes"])
|
||||
if len(chainGroups) > 0 {
|
||||
return chainGroups[len(chainGroups)-1], "chain"
|
||||
}
|
||||
return bestExitDisplayMapSlice(tunnel["inNodeId"]), "entry"
|
||||
}
|
||||
|
||||
func bestExitDisplayMapSlice(v interface{}) []map[string]interface{} {
|
||||
switch arr := v.(type) {
|
||||
case []map[string]interface{}:
|
||||
return arr
|
||||
case []interface{}:
|
||||
out := make([]map[string]interface{}, 0, len(arr))
|
||||
for _, item := range arr {
|
||||
if m, ok := item.(map[string]interface{}); ok {
|
||||
out = append(out, m)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayChainGroups(v interface{}) [][]map[string]interface{} {
|
||||
switch groups := v.(type) {
|
||||
case [][]map[string]interface{}:
|
||||
return groups
|
||||
case []interface{}:
|
||||
out := make([][]map[string]interface{}, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
items := bestExitDisplayMapSlice(group)
|
||||
if len(items) > 0 {
|
||||
out = append(out, items)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayNodeName(source map[string]interface{}, nodeID int64, lookup bestExitNodeNameLookup, fallback string) string {
|
||||
if source != nil {
|
||||
for _, key := range []string{"nodeName", "name"} {
|
||||
if name := strings.TrimSpace(asString(source[key])); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
if lookup != nil {
|
||||
if name, ok := lookup(nodeID); ok && strings.TrimSpace(name) != "" {
|
||||
return strings.TrimSpace(name)
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func bestExitUnknownOwnerName(role string) string {
|
||||
if role == "chain" {
|
||||
return bestExitUnknownChainName
|
||||
}
|
||||
return bestExitUnknownEntryName
|
||||
}
|
||||
|
||||
func (h *Handler) attachBestExitStatesOrLog(items []map[string]interface{}) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
log.Printf("best_exit: attach display state failed: %v", recovered)
|
||||
}
|
||||
}()
|
||||
h.attachBestExitStates(items)
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Replace direct attach calls with panic-safe wrapper**
|
||||
|
||||
Keep `attachBestExitStates` for tests, and use `attachBestExitStatesOrLog` from handlers in Task 3. This step only creates the function above; no handler wiring yet.
|
||||
|
||||
- [ ] **Step 3: Run backend display tests**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run 'TestBestExitDecisionSnapshot|TestBuildBestExitDisplayState' -count=1
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 4: Run gofmt**
|
||||
|
||||
```bash
|
||||
gofmt -w internal/http/handler/tunnel_best_exit_display.go internal/http/handler/tunnel_best_exit_display_test.go
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Commit backend display implementation**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/tunnel_best_exit_display.go go-backend/internal/http/handler/tunnel_best_exit_display_test.go
|
||||
git commit -m "feat: build best exit display state"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Attach Best-Exit State To Tunnel List And Get Responses
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
- Test: `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`
|
||||
|
||||
- [ ] **Step 1: Write failing handler attach tests**
|
||||
|
||||
Append to `go-backend/internal/http/handler/tunnel_best_exit_display_test.go`:
|
||||
|
||||
```go
|
||||
func TestAttachBestExitStatesAddsStateToBestTunnelOnly(t *testing.T) {
|
||||
h := &Handler{bestExit: newBestExitManager()}
|
||||
now := time.Unix(300, 0)
|
||||
h.bestExit.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
|
||||
items := []map[string]interface{}{
|
||||
{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
},
|
||||
{
|
||||
"id": int64(78),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(12)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(40), "strategy": "round"},
|
||||
{"nodeId": int64(41), "strategy": "round"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
h.attachBestExitStates(items)
|
||||
state, ok := items[0]["bestExitState"].(*bestExitDisplayState)
|
||||
if !ok {
|
||||
t.Fatalf("expected bestExitState on best tunnel, got %#v", items[0]["bestExitState"])
|
||||
}
|
||||
if state.Summary != bestExitUnknownExitName || state.Items[0].ExitNodeID != 30 {
|
||||
t.Fatalf("unexpected state with fallback names: %+v", state)
|
||||
}
|
||||
if _, exists := items[1]["bestExitState"]; exists {
|
||||
t.Fatalf("non-best tunnel should not have bestExitState: %+v", items[1])
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run attach test to verify failure**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run TestAttachBestExitStatesAddsStateToBestTunnelOnly -count=1
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Wire tunnel list response**
|
||||
|
||||
In `go-backend/internal/http/handler/handler.go`, change `tunnelList` from:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStatesOrLog(items)
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Wire single tunnel response**
|
||||
|
||||
In `go-backend/internal/http/handler/mutations.go`, change `tunnelGet` from:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
for _, it := range items {
|
||||
if asInt64(it["id"], 0) == id {
|
||||
response.WriteJSON(w, response.OK(it))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```go
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStatesOrLog(items)
|
||||
for _, it := range items {
|
||||
if asInt64(it["id"], 0) == id {
|
||||
response.WriteJSON(w, response.OK(it))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Run focused backend tests**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./internal/http/handler -run 'TestBestExitDecisionSnapshot|TestBuildBestExitDisplayState|TestAttachBestExitStatesAddsStateToBestTunnelOnly' -count=1
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 6: Run gofmt**
|
||||
|
||||
```bash
|
||||
gofmt -w internal/http/handler/handler.go internal/http/handler/mutations.go internal/http/handler/tunnel_best_exit_display.go internal/http/handler/tunnel_best_exit_display_test.go
|
||||
```
|
||||
|
||||
- [ ] **Step 7: Commit response wiring**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/mutations.go go-backend/internal/http/handler/tunnel_best_exit_display.go go-backend/internal/http/handler/tunnel_best_exit_display_test.go
|
||||
git commit -m "feat: expose best exit display state"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Frontend Tunnel List Display
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/pages/tunnel.tsx`
|
||||
|
||||
- [ ] **Step 1: Add TypeScript types**
|
||||
|
||||
In `vite-frontend/src/pages/tunnel.tsx`, add these interfaces after `interface ChainTunnel`:
|
||||
|
||||
```ts
|
||||
interface BestExitStateItem {
|
||||
ownerNodeId: number;
|
||||
ownerNodeName: string;
|
||||
ownerRole: "entry" | "chain";
|
||||
exitNodeId?: number;
|
||||
exitNodeName: string;
|
||||
updatedAt?: number;
|
||||
reason?: string;
|
||||
}
|
||||
|
||||
interface BestExitState {
|
||||
enabled: boolean;
|
||||
summary: string;
|
||||
status: "applied" | "waiting";
|
||||
updatedAt?: number;
|
||||
reason?: string;
|
||||
items: BestExitStateItem[];
|
||||
}
|
||||
```
|
||||
|
||||
Then add the optional field to `interface Tunnel`:
|
||||
|
||||
```ts
|
||||
bestExitState?: BestExitState | null;
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Preserve API state during mapping**
|
||||
|
||||
In `mapTunnelApiItems`, add `bestExitState` to the returned object:
|
||||
|
||||
```ts
|
||||
bestExitState:
|
||||
tunnel.bestExitState && typeof tunnel.bestExitState === "object"
|
||||
? {
|
||||
...tunnel.bestExitState,
|
||||
items: Array.isArray(tunnel.bestExitState.items)
|
||||
? tunnel.bestExitState.items
|
||||
: [],
|
||||
}
|
||||
: null,
|
||||
```
|
||||
|
||||
The mapped object should include this field before `createdTime` or immediately after it.
|
||||
|
||||
- [ ] **Step 3: Add render helpers**
|
||||
|
||||
Add these helper functions after `mapTunnelApiItems` and before `export default function TunnelPage()`:
|
||||
|
||||
```tsx
|
||||
const bestExitOwnerRoleText = (role: BestExitStateItem["ownerRole"]) => {
|
||||
return role === "chain" ? "中转" : "入口";
|
||||
};
|
||||
|
||||
const bestExitDetailTitle = (state?: BestExitState | null) => {
|
||||
if (!state?.enabled || !state.items?.length) {
|
||||
return "";
|
||||
}
|
||||
return state.items
|
||||
.map((item) => {
|
||||
const ownerName = item.ownerNodeName || `${bestExitOwnerRoleText(item.ownerRole)} ${item.ownerNodeId}`;
|
||||
const exitName = item.exitNodeName || "等待探测";
|
||||
return `${ownerName} -> ${exitName}`;
|
||||
})
|
||||
.join("\n");
|
||||
};
|
||||
|
||||
const renderBestExitState = (state?: BestExitState | null) => {
|
||||
if (!state?.enabled) {
|
||||
return null;
|
||||
}
|
||||
const title = bestExitDetailTitle(state);
|
||||
const isWaiting = state.status === "waiting";
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`mt-1 text-[11px] leading-4 ${
|
||||
isWaiting
|
||||
? "text-default-500"
|
||||
: "text-emerald-700 dark:text-emerald-300"
|
||||
}`}
|
||||
title={title || undefined}
|
||||
>
|
||||
最优出口:{state.summary || "等待探测"}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Render in table topology cell**
|
||||
|
||||
In the table topology `<TableCell>` around line 1674, change the cell content from:
|
||||
|
||||
```tsx
|
||||
<div className="flex items-center gap-1.5 text-xs">
|
||||
<span className="font-semibold text-primary-700 dark:text-primary-400">
|
||||
{tunnel.inNodeId?.length || 0}入口
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-secondary-700 dark:text-secondary-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.chainNodes?.length || 0
|
||||
: 0}
|
||||
跳
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-success-700 dark:text-success-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.outNodeId?.length || 0
|
||||
: tunnel.inNodeId?.length || 0}
|
||||
出口
|
||||
</span>
|
||||
</div>
|
||||
```
|
||||
|
||||
to:
|
||||
|
||||
```tsx
|
||||
<div>
|
||||
<div className="flex items-center gap-1.5 text-xs">
|
||||
<span className="font-semibold text-primary-700 dark:text-primary-400">
|
||||
{tunnel.inNodeId?.length || 0}入口
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-secondary-700 dark:text-secondary-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.chainNodes?.length || 0
|
||||
: 0}
|
||||
跳
|
||||
</span>
|
||||
<span className="text-default-400">→</span>
|
||||
<span className="font-semibold text-success-700 dark:text-success-400">
|
||||
{tunnel.type === 2
|
||||
? tunnel.outNodeId?.length || 0
|
||||
: tunnel.inNodeId?.length || 0}
|
||||
出口
|
||||
</span>
|
||||
</div>
|
||||
{renderBestExitState(tunnel.bestExitState)}
|
||||
</div>
|
||||
```
|
||||
|
||||
- [ ] **Step 5: Render in grid card topology section**
|
||||
|
||||
In the grid card topology section, after the closing `</div>` for the topology row at the end of the block containing `出口` and before the enclosing border section closes, add:
|
||||
|
||||
```tsx
|
||||
<div className="text-center">
|
||||
{renderBestExitState(tunnel.bestExitState)}
|
||||
</div>
|
||||
```
|
||||
|
||||
The result should put the best-exit summary under the entry -> hop -> exit row inside the topology section.
|
||||
|
||||
- [ ] **Step 6: Run frontend build**
|
||||
|
||||
Run from `vite-frontend`:
|
||||
|
||||
```bash
|
||||
pnpm run build
|
||||
```
|
||||
|
||||
Expected: PASS with `tsc && vite build` completing successfully.
|
||||
|
||||
- [ ] **Step 7: Commit frontend display**
|
||||
|
||||
```bash
|
||||
git add vite-frontend/src/pages/tunnel.tsx
|
||||
git commit -m "feat: show current best exit in tunnel list"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 5: Full Verification And Review
|
||||
|
||||
**Files:**
|
||||
- Verify only.
|
||||
|
||||
- [ ] **Step 1: Run backend tests**
|
||||
|
||||
Run from `go-backend`:
|
||||
|
||||
```bash
|
||||
go test ./...
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 2: Run frontend build**
|
||||
|
||||
Run from `vite-frontend`:
|
||||
|
||||
```bash
|
||||
pnpm run build
|
||||
```
|
||||
|
||||
Expected: PASS.
|
||||
|
||||
- [ ] **Step 3: Inspect final diff**
|
||||
|
||||
Run from repository root:
|
||||
|
||||
```bash
|
||||
git diff --stat origin/main...HEAD
|
||||
git diff -- go-backend/internal/http/handler/tunnel_best_exit_display.go go-backend/internal/http/handler/tunnel_best_exit_display_test.go go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/mutations.go vite-frontend/src/pages/tunnel.tsx
|
||||
```
|
||||
|
||||
Expected: Diff only adds best-exit display state, response attachment, frontend list display, and tests. It must not change best-exit scoring, switching, runtime chain update, or agent code.
|
||||
|
||||
- [ ] **Step 4: Request final code review**
|
||||
|
||||
Ask a reviewer to check:
|
||||
|
||||
```text
|
||||
Review the best-exit current display implementation. Confirm it only exposes current in-memory best-exit state in tunnel list/get responses and renders it in the tunnel list. Verify it does not change routing, scoring, switching, persistence, or polling behavior.
|
||||
```
|
||||
|
||||
Expected: No blocking findings.
|
||||
|
||||
---
|
||||
|
||||
## Self-Review
|
||||
|
||||
- Spec coverage: Backend response state is Task 2 and Task 3; direct vs final-hop owner semantics are covered by Task 1 tests; frontend list/grid display is Task 4; no polling and no routing changes are preserved by Task 5 review instructions.
|
||||
- Placeholder scan: The plan contains concrete files, function names, code blocks, commands, and expected outcomes.
|
||||
- Type consistency: `BestExitState`, `BestExitStateItem`, `bestExitDisplayState`, `bestExitDisplayItem`, `bestExitDecisionSnapshot`, and `bestExitNodeNameLookup` are defined before use and names match across tasks.
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,133 @@
|
||||
# 允许转发到本地地址开关设计
|
||||
|
||||
**日期**: 2026-04-26
|
||||
**状态**: 待审核
|
||||
**作者**: AI Assistant
|
||||
|
||||
## 概述
|
||||
|
||||
新增一个全局设置开关,控制规则目标地址是否允许指向本地/内网地址。默认关闭,保持当前安全策略不变;开启后,规则创建和编辑时允许将目标地址设置为 `127.0.0.1`、`10.x.x.x`、`172.16-31.x.x`、`192.168.x.x` 等本地或私网地址。
|
||||
|
||||
## 背景
|
||||
|
||||
当前后端在规则创建和编辑时会调用 `IsSafeRemoteAddr()`,统一禁止目标地址指向本地/内网地址,用来降低 SSRF / 开放代理风险。这一行为是全局硬编码的,无法按部署场景调整。
|
||||
|
||||
有些用户需要把规则转发到本机或内网服务,因此需要一个显式、全局的开关来放宽这条限制。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 在设置页提供一个全局开关控制该行为。
|
||||
2. 默认关闭,不改变现有安全默认值。
|
||||
3. 开启后,规则创建和编辑允许本地/内网目标地址。
|
||||
4. 不影响其他安全校验和其他业务流程。
|
||||
|
||||
## 影响范围
|
||||
|
||||
### 后端
|
||||
- `go-backend/internal/http/handler/security_utils.go`
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- `go-backend/internal/http/handler/handler.go`
|
||||
|
||||
### 前端
|
||||
- `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
### 测试
|
||||
- `go-backend/tests/contract/forward_contract_test.go` 或新增独立 contract test
|
||||
|
||||
## 详细设计
|
||||
|
||||
### 1. 配置存储
|
||||
|
||||
使用现有 `vite_config` 表新增一个配置项:
|
||||
|
||||
| name | value | 说明 |
|
||||
|------|-------|------|
|
||||
| `allow_local_remote_addr` | `"1"` / `"0"` | 是否允许规则目标地址指向本地/内网地址 |
|
||||
|
||||
约定:
|
||||
- 未配置时按 `"0"` 处理
|
||||
- `"1"` 表示允许
|
||||
- 其他值一律按关闭处理
|
||||
|
||||
### 2. 后端行为
|
||||
|
||||
新增一个轻量辅助函数,用于读取该配置开关:
|
||||
|
||||
```go
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(cfg.Value) == "1"
|
||||
}
|
||||
```
|
||||
|
||||
在以下路径中应用:
|
||||
- `forwardCreate`
|
||||
- `forwardUpdate`
|
||||
|
||||
行为改为:
|
||||
- 当开关关闭时,继续执行 `IsSafeRemoteAddr(remoteAddr)`
|
||||
- 当开关开启时,跳过这条“本地/内网地址禁止”校验
|
||||
|
||||
这样可以把改动范围限定在规则创建/编辑,不改变其他依赖 `IsSafeRemoteAddr()` 的场景。
|
||||
|
||||
### 3. 前端设置页
|
||||
|
||||
在 `vite-frontend/src/pages/config.tsx` 增加一个全局开关配置项。
|
||||
|
||||
建议文案:
|
||||
|
||||
- 标签:`允许转发到本地地址`
|
||||
- 描述:`开启后,规则目标地址可指向 127.0.0.1、10.x.x.x、172.16-31.x.x、192.168.x.x 等本地或内网地址。默认关闭以降低开放代理风险。`
|
||||
|
||||
控件类型:
|
||||
- 使用现有设置页的布尔开关模式
|
||||
|
||||
默认显示策略:
|
||||
- 不依赖其他配置项
|
||||
- 直接显示在设置页的网络/安全相关区域;若现有页面没有单独分区,则先按现有配置项组织方式加入即可
|
||||
|
||||
### 4. 错误与兼容性
|
||||
|
||||
关闭开关时:
|
||||
- 保持现有错误行为,继续阻止本地/内网地址
|
||||
|
||||
开启开关时:
|
||||
- 仅放开“本地/内网地址禁止”这条限制
|
||||
- 仍保留地址格式解析失败等其他错误
|
||||
|
||||
### 5. 测试
|
||||
|
||||
需要补两类后端契约测试:
|
||||
|
||||
1. 开关关闭时拒绝本地/内网地址
|
||||
- 创建规则时使用本地/内网地址
|
||||
- 断言接口返回非 0 code
|
||||
|
||||
2. 开关开启时允许本地/内网地址
|
||||
- 先写入 `vite_config(name=allow_local_remote_addr, value=1)`
|
||||
- 创建或更新规则时使用相同地址
|
||||
- 断言接口成功
|
||||
|
||||
建议至少覆盖:
|
||||
- create 路径
|
||||
- update 路径
|
||||
- 多目标地址输入(逗号或换行分隔)中包含本地地址时的行为
|
||||
|
||||
## 风险与约束
|
||||
|
||||
1. 该开关会降低默认安全防护,应明确标注风险。
|
||||
2. 这是全局开关,不做用户级或规则级细分控制。
|
||||
3. 该开关只影响规则目标地址校验,不影响其他独立的安全策略。
|
||||
|
||||
## 推荐实施顺序
|
||||
|
||||
1. 先补失败的后端契约测试
|
||||
2. 实现后端配置读取与创建/更新分支控制
|
||||
3. 在设置页增加开关
|
||||
4. 跑后端测试与前端构建验证
|
||||
@@ -0,0 +1,321 @@
|
||||
# 规则每 IP 连接数与限速设计
|
||||
|
||||
**日期**: 2026-04-27
|
||||
**状态**: 待审核
|
||||
**作者**: AI Assistant
|
||||
|
||||
## 概述
|
||||
|
||||
在转发规则的高级设置中新增两类每客户端 IP 限制:每 IP 最大连接数、每 IP 带宽限速。保留现有总量限制语义不变,新增字段只在用户显式配置时生效。
|
||||
|
||||
实现优先复用 GOST 已有能力:`climiters` 的 `$$ N` 表示每个客户端 IP 独立最大连接数;`limiters` 支持 IP/CIDR 级带宽桶,可用 `0.0.0.0/0` 和 `::/0` 实现默认覆盖所有 IPv4/IPv6 客户端的每 IP 带宽限速。
|
||||
|
||||
## 背景
|
||||
|
||||
当前 FLVX 已经支持规则级最大连接数和规则级限速,但这两个限制都是规则总量:
|
||||
|
||||
- `maxConn` 下发为 GOST `climiters` 的 `$ N`,限制整条规则的总并发连接数。
|
||||
- `speedId` 下发为 GOST `limiters` 的 `$ in out`,限制整条规则的总带宽。
|
||||
|
||||
用户需要的是按客户端 IP 隔离的限制,例如每个 IP 最多 5 个连接、每个 IP 最多 10 Mbps,而不是所有客户端共享同一个总量。
|
||||
|
||||
## GOST 能力确认
|
||||
|
||||
### 连接数限制
|
||||
|
||||
`go-gost/x/limiter/conn/conn.go` 已内置以下语义:
|
||||
|
||||
| Key | 含义 |
|
||||
|-----|------|
|
||||
| `$` | 全局连接数限制,所有客户端共享一个 limiter |
|
||||
| `$$` | 每个客户端 IP 独立连接数限制,每个 IP 创建自己的 limiter |
|
||||
| `IP` / `CIDR` | 指定 IP 或 CIDR 的连接数限制 |
|
||||
|
||||
因此每 IP 连接数无需新增 agent 限制器,只需后端下发 `$$ N`。
|
||||
|
||||
### 带宽限制
|
||||
|
||||
`go-gost/x/limiter/traffic/traffic.go` 已内置以下语义:
|
||||
|
||||
| Key | 含义 |
|
||||
|-----|------|
|
||||
| `$` | 服务级总带宽限制 |
|
||||
| `$$` | 连接级带宽限制 |
|
||||
| `IP` / `CIDR` | 客户端 IP 或 CIDR 级带宽限制 |
|
||||
|
||||
CIDR 级限制使用 generator,为命中的客户端 IP 创建独立 limiter。使用 `0.0.0.0/0` 和 `::/0` 可以覆盖所有 IPv4/IPv6 客户端,实现每 IP 带宽限速。
|
||||
|
||||
### 现有缺口
|
||||
|
||||
TCP listener 已在 Accept 后用客户端地址包装连接级 traffic limiter,路径可用于每 IP 带宽。UDP listener 当前只在 PacketConn 上应用服务级 limiter,没有在 `Accept()` 后按客户端 UDP pseudo-connection 包装 limiter,也没有挂接 connection limiter。因此要让 UDP 与 TCP 语义一致,需要补齐 UDP listener 的 per-client wrapper。
|
||||
|
||||
## 目标
|
||||
|
||||
1. 保留现有 `maxConn` 和 `speedId` 的总量语义。
|
||||
2. 在规则上新增每 IP 最大连接数。
|
||||
3. 在规则上新增每 IP 带宽限速。
|
||||
4. 同一规则允许同时配置总量限制和每 IP 限制。
|
||||
5. 普通用户不能设置或修改限速规则字段,保持现有权限模型。
|
||||
6. TCP 和 UDP 入口都尽量遵循相同限制语义。
|
||||
|
||||
## 非目标
|
||||
|
||||
1. 不新增按用户组、节点组、国家地区、ASN 的限制。
|
||||
2. 不新增请求频率限制;本次“每个 IP 限速”指带宽限速,不是新建连接频率。
|
||||
3. 不改变已有 speed limit 规则表的单位和含义。
|
||||
4. 不把用户级默认最大连接数改成每 IP 语义;用户级 `maxConn` 继续作为默认总连接数。
|
||||
|
||||
## 数据模型
|
||||
|
||||
在 `forward` 表新增两个字段:
|
||||
|
||||
| 字段 | 类型 | 默认 | 说明 |
|
||||
|------|------|------|------|
|
||||
| `ip_max_conn` | int | `0` | 每 IP 最大连接数,`0` 表示不启用 |
|
||||
| `ip_speed_id` | nullable int64 | `NULL` | 每 IP 带宽限速规则 ID,`NULL` 表示不启用 |
|
||||
|
||||
Go 模型新增:
|
||||
|
||||
```go
|
||||
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
|
||||
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
|
||||
```
|
||||
|
||||
字段会通过现有 auto-migrate 机制创建,保持 SQLite/PostgreSQL 兼容,不使用 SQLite 不兼容的 GORM tags。
|
||||
|
||||
## API 行为
|
||||
|
||||
### 创建规则
|
||||
|
||||
`/forward/create` 新增入参:
|
||||
|
||||
```json
|
||||
{
|
||||
"ipMaxConn": 5,
|
||||
"ipSpeedId": 123
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- `ipMaxConn` 缺省或小于等于 `0` 时按 `0` 存储,不启用每 IP 连接数限制。
|
||||
- `ipSpeedId` 缺省或不存在时存为 `NULL`,不启用每 IP 带宽限速。
|
||||
- `ipSpeedId` 指向不存在的限速规则时按 `NULL` 处理,沿用现有 `speedId` 的容错策略。
|
||||
- 普通用户提交非空 `ipSpeedId` 时返回错误,保持与 `speedId` 一致的权限边界。
|
||||
|
||||
### 更新规则
|
||||
|
||||
`/forward/update` 新增入参:
|
||||
|
||||
```json
|
||||
{
|
||||
"ipMaxConn": 5,
|
||||
"ipSpeedId": 123
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- 未提交 `ipMaxConn` 时保留原值;提交空值或 `0` 时清除每 IP 连接数限制。
|
||||
- 未提交 `ipSpeedId` 时保留原值;提交 `null` 时清除每 IP 带宽限速。
|
||||
- 普通用户不能把 `ipSpeedId` 改成不同的非空值。
|
||||
- 更新后重新同步运行时服务和 limiter。
|
||||
|
||||
### 列表返回
|
||||
|
||||
`/forward/list` 返回项新增:
|
||||
|
||||
```json
|
||||
{
|
||||
"ipMaxConn": 5,
|
||||
"ipSpeedId": 123,
|
||||
"ipSpeedLimitName": "每IP 10Mbps"
|
||||
}
|
||||
```
|
||||
|
||||
`ipSpeedLimitName` 可选,但建议返回,便于前端显示缺失或已删除的限速规则。
|
||||
|
||||
## 后端运行时同步
|
||||
|
||||
### 连接数限制器
|
||||
|
||||
将现有连接限制器构建从单一总量扩展为组合规则。
|
||||
|
||||
当前行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "rule_conn_limit_42",
|
||||
"limits": ["$ 100"]
|
||||
}
|
||||
```
|
||||
|
||||
新增行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "rule_conn_limit_42",
|
||||
"limits": ["$ 100", "$$ 5"]
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- `maxConn > 0` 时追加 `$ maxConn`。
|
||||
- `ipMaxConn > 0` 时追加 `$$ ipMaxConn`。
|
||||
- 如果规则未配置 `maxConn` 且用户有 `MaxConn > 0`,继续继承用户级总连接数,追加 `$ user.MaxConn`。
|
||||
- 如果两者都没有,则不下发 `climiter`,服务不引用 `climiter`。
|
||||
- limiter 名称继续优先使用 `rule_conn_limit_<forwardID>`;只有用户级默认总连接数且规则没有任何连接限制时可继续使用 `user_conn_limit_<userID>`,避免不必要的 per-rule limiter。
|
||||
|
||||
### 带宽限制器
|
||||
|
||||
将现有规则限速从单一 `speedId` 扩展为组合 limiter。
|
||||
|
||||
当前行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "123",
|
||||
"limits": ["$ 1.3MB 1.3MB"]
|
||||
}
|
||||
```
|
||||
|
||||
新增每 IP 行为:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "rule_traffic_limit_42",
|
||||
"limits": [
|
||||
"$ 1.3MB 1.3MB",
|
||||
"0.0.0.0/0 1.3MB 1.3MB",
|
||||
"::/0 1.3MB 1.3MB"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- 只有总量 `speedId` 时,保持现有名称和下发路径,服务继续引用 `speedId` 字符串。
|
||||
- 只有每 IP `ipSpeedId` 时,创建 `rule_traffic_limit_<forwardID>`,只包含 IPv4/IPv6 CIDR 行。
|
||||
- 总量和每 IP 同时存在时,创建 `rule_traffic_limit_<forwardID>`,同时包含 `$` 和 CIDR 行。
|
||||
- 如果规则没有 `speedId`,则总量仍可继承 user tunnel 的 `speedId`,保持现有 fallback 语义;当继承的总量限速与 `ipSpeedId` 同时存在时,也使用 `rule_traffic_limit_<forwardID>` 组合 limiter。
|
||||
- 每 IP 限速不从 user tunnel 继承,只由规则字段控制。
|
||||
- `AddLimiters` 失败且提示已存在时,使用 `UpdateLimiters` 更新。
|
||||
|
||||
### 服务配置
|
||||
|
||||
`buildForwardServiceConfigs` 需要从当前 `limiterID *int64` / `cLimiterName string` 扩展为更明确的运行时限制描述,例如:
|
||||
|
||||
```go
|
||||
type forwardRuntimeLimiters struct {
|
||||
TrafficLimiter string
|
||||
ConnLimiter string
|
||||
}
|
||||
```
|
||||
|
||||
服务配置只关心最终引用的 limiter 名称:
|
||||
|
||||
- `service["limiter"] = runtimeLimiters.TrafficLimiter`
|
||||
- `service["climiter"] = runtimeLimiters.ConnLimiter`
|
||||
|
||||
这样可以把“如何构建 limiter payload”的逻辑和“如何构建 service JSON”的逻辑分开。
|
||||
|
||||
## Agent/GOST 调整
|
||||
|
||||
### WebSocket 命令
|
||||
|
||||
当前 agent WebSocket 已支持:
|
||||
|
||||
- `AddLimiters` / `UpdateLimiters` / `DeleteLimiters`
|
||||
- `AddCLimiters` / `UpdateCLimiters` / `DeleteCLimiters`
|
||||
|
||||
本设计无需新增命令类型。
|
||||
|
||||
### UDP listener
|
||||
|
||||
补齐 `go-gost/x/listener/udp/listener.go` 的 `Accept()` 包装逻辑,使 UDP pseudo-connection 与 TCP listener 一致:
|
||||
|
||||
- 对 `l.options.ConnLimiter` 按客户端地址应用连接数限制。
|
||||
- 对 `l.options.TrafficLimiter` 按 `conn.RemoteAddr().String()` 应用连接级 traffic wrapper。
|
||||
|
||||
需要注意 UDP pseudo-connection 的生命周期由内部 UDP listener 的 TTL/keepalive 控制;connection limiter 必须在 pseudo-connection 关闭时释放计数。
|
||||
|
||||
## 前端设计
|
||||
|
||||
在 `vite-frontend/src/pages/forward.tsx` 的规则高级设置中新增两个控件:
|
||||
|
||||
1. `每 IP 最大连接数`
|
||||
- 类型:number input。
|
||||
- 文案:`每个客户端 IP 可同时建立的最大连接数;0 或空表示不限制。`
|
||||
- 字段:`ipMaxConn`。
|
||||
|
||||
2. `每 IP 限速`
|
||||
- 类型:Select,复用现有限速规则列表。
|
||||
- 文案:`每个客户端 IP 独享该带宽限制;不选择表示不限制。`
|
||||
- 字段:`ipSpeedId`。
|
||||
- 只对管理员显示,保持与 `规则限速` 一致。
|
||||
|
||||
前端类型需要同步更新:
|
||||
|
||||
- `ForwardApiItem`
|
||||
- `ForwardMutationPayload`
|
||||
- `ForwardForm` 或页面内等价类型
|
||||
|
||||
## 错误处理与兼容性
|
||||
|
||||
1. 旧数据默认 `ip_max_conn=0`、`ip_speed_id=NULL`,行为与当前版本一致。
|
||||
2. 现有 agent 已支持 limiter 命令和 GOST limiter 语法;发布时需要包含 UDP 修复,才能让 TCP/UDP 都获得完整语义。
|
||||
3. 节点离线时沿用现有 warning 行为,规则仍可保存,在线节点跳过下发。
|
||||
4. 如果每 IP speed limit ID 被删除,更新时按 `NULL` 处理,列表页可提示或自动清除,和现有 `speedId` 行为一致。
|
||||
5. 如果 IPv6 CIDR 在某些监听路径未命中,IPv4 行仍正常生效;测试应覆盖 IPv4,IPv6 通过 payload 合同保证下发。
|
||||
|
||||
## 测试计划
|
||||
|
||||
### 后端 contract 测试
|
||||
|
||||
新增或扩展 `go-backend/tests/contract/max_conn_limit_contract_test.go`:
|
||||
|
||||
1. 创建规则时设置 `ipMaxConn=5`,断言 `AddCLimiters` payload 包含 `$$ 5`。
|
||||
2. 同时设置 `maxConn=100` 和 `ipMaxConn=5`,断言 payload 包含 `$ 100` 和 `$$ 5`。
|
||||
3. 用户级 `MaxConn` 存在且规则 `ipMaxConn=5` 时,断言 payload 包含 `$ userMaxConn` 和 `$$ 5`。
|
||||
|
||||
新增每 IP 限速 contract 测试:
|
||||
|
||||
1. 创建规则时设置 `ipSpeedId`,断言 `AddLimiters` payload 包含 `0.0.0.0/0 ...` 和 `::/0 ...`。
|
||||
2. 同时设置 `speedId` 和 `ipSpeedId`,断言组合 limiter 包含 `$ ...` 与两个 CIDR 行,服务引用 `rule_traffic_limit_<forwardID>`。
|
||||
3. 普通用户提交 `ipSpeedId` 返回错误。
|
||||
|
||||
### Repository/API 测试
|
||||
|
||||
1. `CreateForwardTx`、`UpdateForward`、列表查询读写 `ip_max_conn` 和 `ip_speed_id`。
|
||||
2. `/forward/list` 返回 `ipMaxConn`、`ipSpeedId`。
|
||||
|
||||
### GOST/x 测试
|
||||
|
||||
1. `go-gost/x/limiter/conn`:验证 `$$ N` 为不同 IP 创建独立 limiter。
|
||||
2. `go-gost/x/limiter/traffic`:验证 `0.0.0.0/0` 为不同 IPv4 创建独立 limiter。
|
||||
3. UDP listener:验证 Accept 返回的 UDP pseudo-connection 关闭后释放 connection limiter。
|
||||
|
||||
### 验证命令
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd go-gost/x && go test ./limiter/... ./listener/udp/...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## 推荐实施顺序
|
||||
|
||||
1. 后端模型、repo DTO、API 字段读写。
|
||||
2. 后端 limiter payload 构建与服务引用重构。
|
||||
3. Contract 测试覆盖连接数和带宽 payload。
|
||||
4. GOST UDP listener per-client wrapper 与相关测试。
|
||||
5. 前端高级设置表单和类型更新。
|
||||
6. 运行后端测试、GOST/x 相关测试、前端构建。
|
||||
|
||||
## 风险
|
||||
|
||||
1. UDP pseudo-connection 生命周期和 TCP 连接不同,连接数释放必须依赖 Close 包装正确执行。
|
||||
2. 总带宽和每 IP 带宽组合时 limiter 名称从纯 speed ID 变为 rule-level 名称,需要确保更新已有规则时不会留下错误引用。
|
||||
3. 旧节点如果没有 UDP wrapper 修复,TCP 生效但 UDP 每 IP 语义可能不完整;发布时应要求 agent 同步升级。
|
||||
4. 每 IP 带宽是每个入口节点本地独立限制,不是跨节点全局聚合限制。
|
||||
@@ -0,0 +1,74 @@
|
||||
# Monitoring Retention And Storage Display Design
|
||||
|
||||
## Goal
|
||||
|
||||
Add an administrator-facing configuration for monitoring data retention and display the current database storage usage in the configuration page.
|
||||
|
||||
## Scope
|
||||
|
||||
- Add a single config key: `monitor_retention_days`.
|
||||
- Default retention is `7` days.
|
||||
- Apply the retention window uniformly to:
|
||||
- `node_metric`
|
||||
- `tunnel_metric`
|
||||
- `service_monitor_result`
|
||||
- `tunnel_quality`
|
||||
- Show database usage on the config page as a read-only operational value.
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- No per-table retention settings.
|
||||
- No manual purge button.
|
||||
- No database vacuum/compaction action.
|
||||
- No frontend test framework changes.
|
||||
|
||||
## Backend Design
|
||||
|
||||
### Retention Config
|
||||
|
||||
- Store `monitor_retention_days` in `vite_config`, consistent with existing site settings.
|
||||
- Accept integer values from `1` through `3650`.
|
||||
- Missing or invalid stored values fall back to `7` days.
|
||||
- `normalizeAndValidateConfigValue` rejects invalid user-submitted values so bad config does not get saved through the API.
|
||||
|
||||
### Cleanup Flow
|
||||
|
||||
- `metrics.IngestionService.pruneMetrics()` reads `monitor_retention_days` from the repository each hourly cleanup cycle.
|
||||
- The computed cutoff is used for `node_metric`, `tunnel_metric`, and `service_monitor_result`.
|
||||
- `tunnel_quality` uses the same retention config.
|
||||
- `tunnel_quality` cleanup must run even when real-time tunnel quality probing is disabled; disabling probing should stop new probe writes, not stop cleanup.
|
||||
|
||||
### Database Storage API
|
||||
|
||||
- Add an admin-only API endpoint for storage summary, for example `/api/v1/system/storage`.
|
||||
- Response fields:
|
||||
- `dbType`: `sqlite` or `postgres`
|
||||
- `databaseSizeBytes`: raw byte count
|
||||
- `databaseSizeText`: human-readable formatted size
|
||||
- SQLite implementation reports the DB file size and includes `-wal` and `-shm` sidecar files when present.
|
||||
- PostgreSQL implementation uses `pg_database_size(current_database())`.
|
||||
- If size cannot be determined, return an API error rather than a misleading zero.
|
||||
|
||||
## Frontend Design
|
||||
|
||||
- Add `monitor_retention_days` to the config page.
|
||||
- Label: `监控数据保留天数`.
|
||||
- Description: `统一清理节点指标、隧道流量、服务监控结果和隧道质量历史;默认 7 天。`
|
||||
- Use a regular numeric input through the existing config rendering path.
|
||||
- Fetch database storage summary when the config page loads.
|
||||
- Display a read-only card/row named `数据库占用` with `databaseSizeText`.
|
||||
- If fetching fails, show `获取失败` and keep config editing usable.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- Invalid retention values return a validation error on save.
|
||||
- Cleanup logs individual prune failures and continues with other tables, matching existing monitoring cleanup behavior.
|
||||
- Storage summary failures are non-blocking in the frontend.
|
||||
|
||||
## Testing
|
||||
|
||||
- Backend unit tests for retention config parsing and validation.
|
||||
- Backend tests proving custom retention is used by monitoring cleanup.
|
||||
- Backend API/repository test for SQLite storage size returning a non-negative byte count and formatted text.
|
||||
- Run `go test ./...` in `go-backend`.
|
||||
- Run `pnpm run build` in `vite-frontend`.
|
||||
@@ -0,0 +1,156 @@
|
||||
# Best Exit Current Selection Display Design
|
||||
|
||||
## Goal
|
||||
|
||||
When a tunnel uses the `best` multi-exit strategy, show the currently applied best exit in the tunnel list information. Users should be able to see which exit is currently selected without opening logs or diagnosing the tunnel manually.
|
||||
|
||||
The display is informational only. It must not change routing, scoring, switching behavior, or the saved tunnel configuration.
|
||||
|
||||
## Current Context
|
||||
|
||||
- `3.0.0-beta6` adds `best` as a multi-exit strategy.
|
||||
- Runtime selection is stored in the backend `bestExitManager` in memory, keyed by `TunnelID + OwnerNodeID`.
|
||||
- Direct multi-entry tunnels make one independent best-exit decision per entry node.
|
||||
- Tunnels with intermediate chain hops make one independent best-exit decision per final-hop chain node before the exits.
|
||||
- `tunnelList` and `tunnelGet` currently return `repo.ListTunnels()` output directly, so frontend tunnel data only includes configured exits from the database, not the currently applied runtime choice.
|
||||
- The frontend tunnel page maps API items in `vite-frontend/src/pages/tunnel.tsx` and renders list information from that data.
|
||||
|
||||
## User Decisions
|
||||
|
||||
- Show the current best-exit choice in the tunnel list information.
|
||||
- Use a summary plus detail model for multiple owners.
|
||||
- Follow the existing tunnel list refresh cadence; do not add polling or a realtime stream in this phase.
|
||||
- Work text-only; no visual companion is needed.
|
||||
|
||||
## Approach
|
||||
|
||||
Extend the existing tunnel list/detail response with a lightweight runtime state object for `best` tunnels, then render that state beside the tunnel's exit/strategy information in the existing frontend list UI.
|
||||
|
||||
This keeps the display close to the data users already inspect and avoids a separate API or extra frontend request.
|
||||
|
||||
## Backend Design
|
||||
|
||||
### Response Shape
|
||||
|
||||
Add a `bestExitState` object to each tunnel item returned by `tunnelList` and `tunnelGet` when the tunnel has a multi-exit group whose strategy is `best`.
|
||||
|
||||
Response shape:
|
||||
|
||||
```json
|
||||
{
|
||||
"enabled": true,
|
||||
"summary": "香港节点",
|
||||
"status": "applied",
|
||||
"updatedAt": 1777584000000,
|
||||
"reason": "current exit remains best",
|
||||
"items": [
|
||||
{
|
||||
"ownerNodeId": 10,
|
||||
"ownerNodeName": "入口 A",
|
||||
"ownerRole": "entry",
|
||||
"exitNodeId": 30,
|
||||
"exitNodeName": "香港节点",
|
||||
"updatedAt": 1777584000000,
|
||||
"reason": "current exit remains best"
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
If the tunnel is not using `best`, omit `bestExitState` or set it to `null`.
|
||||
|
||||
### Owner Semantics
|
||||
|
||||
The display must match the routing model:
|
||||
|
||||
- If there are no middle chain hops, each entry node is an owner.
|
||||
- If there are middle chain hops, each node in the final middle-hop group is an owner.
|
||||
|
||||
Each owner can have a different current best exit. The UI must not imply that a multi-owner tunnel has one global best exit when the owners differ.
|
||||
|
||||
### Summary Rules
|
||||
|
||||
- If all owners currently apply the same exit, `summary` is that exit node name.
|
||||
- If owners apply different exits, `summary` is `多个出口`.
|
||||
- If no applied decision exists yet, `summary` is `等待探测`.
|
||||
- If the tunnel has only one exit, `bestExitState` is not needed because there is no dynamic choice.
|
||||
|
||||
### State Source
|
||||
|
||||
Use the in-memory `bestExitManager` as the source of currently applied decisions.
|
||||
|
||||
Add a read-only snapshot method that returns defensive copies of decision state without exposing mutable internal slices. The handler should convert node IDs to display names from the existing tunnel response data first, then fall back to `h.getNodeRecord` only when the current response does not contain the node.
|
||||
|
||||
The feature should not persist current choices to the database in this phase. A panel restart may reset the displayed runtime state to `等待探测` until the prober initializes it again from the current saved first exit.
|
||||
|
||||
## Frontend Design
|
||||
|
||||
Extend the tunnel item type with optional `bestExitState`.
|
||||
|
||||
In the tunnel list, only render the current best-exit display when:
|
||||
|
||||
- `bestExitState.enabled === true`, or
|
||||
- the tunnel has an exit group with `strategy === "best"` and the backend returns a waiting state.
|
||||
|
||||
Display format:
|
||||
|
||||
- Single applied exit: `最优出口:香港节点`
|
||||
- Multiple applied exits: `最优出口:多个出口`
|
||||
- Waiting: `最优出口:等待探测`
|
||||
|
||||
For multiple owners, render the summary as compact secondary text in the topology/list information cell and set its native `title` attribute to newline-separated detail rows. This avoids adding a new UI dependency or a custom popover. Detail rows should use:
|
||||
|
||||
```text
|
||||
入口 A -> 香港节点
|
||||
入口 B -> 日本节点
|
||||
```
|
||||
|
||||
For tunnels with middle chain hops, label owners as chain nodes when useful:
|
||||
|
||||
```text
|
||||
中转 M1 -> 香港节点
|
||||
中转 M2 -> 日本节点
|
||||
```
|
||||
|
||||
Do not add a new periodic refresh. The display updates when the existing tunnel list is refreshed.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- If the manager has no decision for an owner, show that owner as `等待探测`.
|
||||
- If an exit node ID no longer exists in the current tunnel response, show `未知出口` for that item and keep the list usable.
|
||||
- If an owner node ID no longer exists, show `未知入口` or `未知中转` based on the owner role.
|
||||
- If the backend cannot compute state for one tunnel, omit `bestExitState` for that tunnel and log the error; do not fail the whole tunnel list response.
|
||||
|
||||
## Testing
|
||||
|
||||
Backend tests:
|
||||
|
||||
- `bestExitManager` snapshot returns applied exit IDs without exposing mutable manager state.
|
||||
- Direct multi-entry `best` tunnel produces one display item per entry owner.
|
||||
- Middle-hop tunnel produces one display item per final-hop owner.
|
||||
- Summary is the single exit name when all owners choose the same exit.
|
||||
- Summary is `多个出口` when owners choose different exits.
|
||||
- Summary is `等待探测` when no applied decision exists.
|
||||
- Non-`best` tunnels do not receive `bestExitState`.
|
||||
|
||||
Frontend verification:
|
||||
|
||||
- Tunnel list renders `最优出口:<name>` for a single applied exit.
|
||||
- Tunnel list renders `最优出口:多个出口` plus owner details for multiple applied exits.
|
||||
- Tunnel list renders `最优出口:等待探测` for waiting state.
|
||||
- `pnpm run build` passes.
|
||||
|
||||
Verification commands:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- Do not add a new realtime stream or polling loop.
|
||||
- Do not add a detailed best-exit scoring dashboard.
|
||||
- Do not persist current best-exit choices to the database.
|
||||
- Do not change switching thresholds, probing targets, or runtime chain update behavior.
|
||||
- Do not change existing non-`best` tunnel display behavior.
|
||||
@@ -0,0 +1,186 @@
|
||||
# Best Exit Selection Design
|
||||
|
||||
## Goal
|
||||
|
||||
Add a multi-exit tunnel strategy named `best` that always sends new connections through the currently best-quality exit. The feature should prevent traffic from continuing to use an exit whose latency or packet loss has degraded while the exit is still technically online.
|
||||
|
||||
Existing connections must not be interrupted. Switching affects only new connections created after the runtime chain update is applied.
|
||||
|
||||
## Current Context
|
||||
|
||||
- Tunnel forwarding stores entry, chain, and exit nodes in `chain_tunnel`.
|
||||
- Multi-exit runtime chains are currently rendered as one GOST hop with multiple nodes.
|
||||
- GOST selectors support `fifo`, `round`, `rand`, and `hash`, plus fail filtering through `maxFails` and `failTimeout`.
|
||||
- The current fail filter only reacts to dial, handshake, or transport failures. It does not react to high latency when the exit is still reachable.
|
||||
- `tunnel_quality_prober` already runs panel-side TCP probes and stores tunnel quality history, but it currently probes representative nodes and does not drive runtime routing decisions.
|
||||
|
||||
## User Decisions
|
||||
|
||||
- Add a `best` option for multi-exit tunnels.
|
||||
- `best` means always choose the current best exit for new connections.
|
||||
- Score exits by end-to-end quality.
|
||||
- Keep the existing public probe target: `www.bing.com:443`.
|
||||
- Do not disrupt established connections.
|
||||
|
||||
## Approach
|
||||
|
||||
Implement `best` as a panel-driven control-plane strategy.
|
||||
|
||||
The database stores the user's intended strategy as `best`. When the panel renders runtime GOST config for a `best` exit group, it sends a GOST selector strategy of `fifo`. The panel dynamically sorts the candidate exits so the current best exit is first. GOST then chooses the first node for new connections.
|
||||
|
||||
This avoids adding active probing logic inside every GOST agent and reuses the existing panel-to-agent command path.
|
||||
|
||||
## Components
|
||||
|
||||
### Frontend
|
||||
|
||||
The tunnel form adds `最优` to the multi-exit load strategy selector.
|
||||
|
||||
- Label: `最优`
|
||||
- Value: `best`
|
||||
- Scope: tunnel forwarding exit groups, alongside `主备/fifo`, `轮询/round`, and `随机/rand`
|
||||
- Create and edit forms must submit and restore `best` unchanged.
|
||||
|
||||
### Backend Data Model
|
||||
|
||||
No schema change is required.
|
||||
|
||||
The existing `chain_tunnel.strategy` column stores `best`. Repository and handler paths should preserve the value in API responses and updates.
|
||||
|
||||
### Runtime Chain Rendering
|
||||
|
||||
When building runtime chain config:
|
||||
|
||||
- If the configured strategy is not `best`, keep existing behavior.
|
||||
- If the configured strategy is `best`, emit GOST selector strategy `fifo`.
|
||||
- Sort the target nodes using the panel's latest best-exit decision before rendering the node list.
|
||||
- If no quality decision exists yet, keep the saved node order.
|
||||
|
||||
This preserves the user's `best` intent in storage while using a GOST selector that can execute the panel's sorted decision.
|
||||
|
||||
### Quality Prober
|
||||
|
||||
Extend `tunnel_quality_prober` to evaluate all candidates in `best` exit groups.
|
||||
|
||||
For each chain owner node and candidate exit, measure:
|
||||
|
||||
- Chain owner node to candidate exit using TCP ping.
|
||||
- Candidate exit to `www.bing.com:443` using TCP ping.
|
||||
|
||||
For direct entry-to-exit tunnels, each entry node owns its own chain decision. For tunnels with intermediate chain hops, each node in the last hop group before the exits owns its own chain decision. This allows different entry or chain nodes to choose different best exits when their path quality differs.
|
||||
|
||||
### Scoring
|
||||
|
||||
Each exit candidate gets an end-to-end score for a specific chain owner node.
|
||||
|
||||
- Total latency is the sum of owner-to-exit latency and exit-to-Bing latency.
|
||||
- Total loss combines both legs by success probability: `1 - (1 - lossA) * (1 - lossB)`.
|
||||
- Failed or unreachable candidates are sorted behind successful candidates.
|
||||
- The score should heavily penalize packet loss so that low-latency but lossy exits are not selected over stable exits.
|
||||
|
||||
A practical scoring formula can be:
|
||||
|
||||
```text
|
||||
score = totalLatencyMs + (totalLossPercent * lossPenaltyMsPerPercent)
|
||||
```
|
||||
|
||||
Use `lossPenaltyMsPerPercent = 100` initially. For example, 5% loss adds 500ms to the score.
|
||||
|
||||
### Switching Rules
|
||||
|
||||
The panel should not update chains on every probe round.
|
||||
|
||||
Switch only when all conditions are true:
|
||||
|
||||
- The candidate best exit is different from the currently applied first exit.
|
||||
- The candidate is successful.
|
||||
- The candidate remains best for consecutive probe rounds.
|
||||
- The candidate beats the current exit by a minimum advantage threshold.
|
||||
- The chain owner node has passed a minimum switch cooldown.
|
||||
|
||||
Initial constants:
|
||||
|
||||
- Consecutive confirmations: 3 rounds.
|
||||
- Switch cooldown: 30 seconds per chain owner node.
|
||||
- Minimum advantage: the candidate score must improve by at least `max(20ms, currentScore * 0.15)`.
|
||||
|
||||
If all exits fail, keep the current runtime order and do not issue a destructive update.
|
||||
|
||||
### Runtime Update
|
||||
|
||||
When a `best` chain owner node changes best exit:
|
||||
|
||||
1. Rebuild that node's `chains_<tunnelID>` payload with the best exit first and remaining candidates sorted by quality for that node.
|
||||
2. Send `UpdateChains` to that chain owner node.
|
||||
3. Do not restart or update tunnel services.
|
||||
4. Record success or failure in logs and in the in-memory decision state.
|
||||
|
||||
This affects only future connections. Existing TCP connections keep using the `net.Conn` created before the update and continue through their original exit.
|
||||
|
||||
### Agent Safety Improvement
|
||||
|
||||
The current agent `UpdateChains` path unregisters the old chain before registering the new chain. This does not kill existing connections, but it creates a small window where a new connection can fail because the chain name is temporarily absent.
|
||||
|
||||
Improve the update path so it parses the new chain first and only replaces the registered chain after parsing succeeds. The replacement window should be as small as possible. If parsing fails, the old chain must remain active.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- If probing one candidate fails, continue scoring other candidates.
|
||||
- If a chain owner node is offline or times out, skip decisions for that owner during the round instead of marking every candidate failed.
|
||||
- If a candidate has no successful required probe data, mark it failed for that round.
|
||||
- If `UpdateChains` fails, keep the current applied order and retry on a later round.
|
||||
- If the tunnel has one exit or an incomplete config, `best` behaves like the saved order and does not trigger dynamic switching.
|
||||
- If `monitor_tunnel_quality_enabled=false`, dynamic `best` switching pauses. The last applied runtime order remains in effect.
|
||||
|
||||
## Observability
|
||||
|
||||
The prober should maintain in-memory decision state per `best` tunnel and chain owner node.
|
||||
|
||||
Useful fields:
|
||||
|
||||
- Tunnel ID and chain owner node ID.
|
||||
- Current applied best exit node ID.
|
||||
- Candidate best exit node ID.
|
||||
- Candidate scores.
|
||||
- Last switch timestamp.
|
||||
- Last switch result.
|
||||
- Reason for not switching, such as cooldown, insufficient advantage, candidate unstable, or all exits failed.
|
||||
|
||||
Initial UI scope is limited to supporting create, update, and display of the `best` strategy. A later enhancement can expose current best exit and candidate scores in the tunnel monitor view.
|
||||
|
||||
## Testing
|
||||
|
||||
Backend tests:
|
||||
|
||||
- Score calculation orders candidates by latency and packet loss.
|
||||
- Packet loss penalty prevents lossy exits from winning only because latency is low.
|
||||
- All-failed candidates do not trigger a switch.
|
||||
- Consecutive confirmation and cooldown prevent flapping.
|
||||
- `strategy=best` persists in `chain_tunnel.strategy` and is returned by tunnel list/get APIs.
|
||||
- Runtime rendering maps `best` to GOST `fifo` and places the chosen best exit first.
|
||||
|
||||
Agent tests:
|
||||
|
||||
- `UpdateChains` parse failure keeps the old chain registered.
|
||||
- Successful `UpdateChains` updates the chain used by new connections.
|
||||
|
||||
Frontend verification:
|
||||
|
||||
- Tunnel form includes `最优` in the exit strategy selector.
|
||||
- Existing tunnels with `strategy=best` render correctly.
|
||||
- Create and update requests submit `best` unchanged.
|
||||
|
||||
Verification commands:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
(cd go-gost && go test ./...)
|
||||
(cd vite-frontend && pnpm run build)
|
||||
```
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- Do not move existing live connections to a new exit.
|
||||
- Do not add per-tunnel custom probe targets in this phase.
|
||||
- Do not implement active best-exit probing inside GOST agents.
|
||||
- Do not add a detailed best-exit UI dashboard in this phase.
|
||||
@@ -0,0 +1,180 @@
|
||||
# Custom Best-Exit Probe Target Design
|
||||
|
||||
Date: 2026-05-01
|
||||
Status: Approved design
|
||||
|
||||
## Goal
|
||||
|
||||
Allow each tunnel to define the TCP target used for exit-side quality probing instead of always probing `www.bing.com:443`.
|
||||
|
||||
The custom target must be used consistently by:
|
||||
|
||||
- `best` exit scoring: each exit probes the configured target to measure exit-to-public quality.
|
||||
- Tunnel quality monitoring: the existing exit-side quality check probes the same configured target.
|
||||
|
||||
If a tunnel does not configure a target, behavior remains compatible with today: `www.bing.com:443`.
|
||||
|
||||
## Non-Goals
|
||||
|
||||
- Do not add HTTP/HTTPS request probing in this phase. The probe remains TCP host/port measurement.
|
||||
- Do not add a global default target setting in this phase.
|
||||
- Do not require existing tunnels to be edited or migrated manually.
|
||||
- Do not change the `best` switching thresholds, confirmation rounds, cooldowns, or runtime chain ordering semantics.
|
||||
- Do not add frontend test infrastructure.
|
||||
|
||||
## User-Facing Behavior
|
||||
|
||||
Each tunnel form gets a compact quality target section:
|
||||
|
||||
- Host input, placeholder `www.bing.com`.
|
||||
- Port input, placeholder `443`.
|
||||
- Helper text: this target is used for tunnel quality detection and `best` optimal-exit scoring; leaving it empty uses `www.bing.com:443`.
|
||||
|
||||
Tunnel list/get responses include the configured target so edit forms can round-trip it. The UI displays the effective target near quality/best-exit information as `测试目标:host:port`.
|
||||
|
||||
## Data Model
|
||||
|
||||
Add nullable/default-compatible fields to `model.Tunnel`:
|
||||
|
||||
- `ProbeTargetHost string` mapped to `probe_target_host`, `type:text`, default `''`.
|
||||
- `ProbeTargetPort int` mapped to `probe_target_port`, default `0`.
|
||||
|
||||
Effective target resolution:
|
||||
|
||||
- If `ProbeTargetHost` is non-empty and `ProbeTargetPort` is valid, use it.
|
||||
- Otherwise use `www.bing.com:443`.
|
||||
|
||||
The existing `TunnelQuality` persisted fields `exit_to_bing_latency` and `exit_to_bing_loss` remain unchanged for compatibility. They will semantically mean exit-to-configured-test-target after this change. API/UI labels should avoid saying `Bing` for new displays.
|
||||
|
||||
## Validation
|
||||
|
||||
On create/update:
|
||||
|
||||
- Empty host and empty/zero port are allowed and mean default target.
|
||||
- If either host or port is set, validate both as a pair.
|
||||
- Host is trimmed and must not contain URL scheme, path, query, or whitespace.
|
||||
- Host can be a domain, IPv4, or IPv6 literal. Bracketed IPv6 input should be normalized by removing surrounding brackets.
|
||||
- Port must be an integer from `1` to `65535`.
|
||||
- Do not perform network probing during save; external network failures must not block configuration changes.
|
||||
|
||||
Errors should be specific, for example:
|
||||
|
||||
- `测试目标 Host 不能为空`
|
||||
- `测试目标端口必须是 1-65535`
|
||||
- `测试目标 Host 不能包含协议或路径`
|
||||
|
||||
## Backend Flow
|
||||
|
||||
Introduce a small value/helper near the tunnel quality and best-exit code:
|
||||
|
||||
```go
|
||||
type tunnelProbeTarget struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
```
|
||||
|
||||
Helpers:
|
||||
|
||||
- `defaultTunnelProbeTarget() tunnelProbeTarget` returns `www.bing.com:443`.
|
||||
- `normalizeTunnelProbeTarget(host string, port int) (tunnelProbeTarget, bool, error)` validates user input; the boolean indicates whether the user explicitly configured a target.
|
||||
- `effectiveTunnelProbeTarget(tunnel *model.Tunnel) tunnelProbeTarget` returns configured target or default.
|
||||
|
||||
Use the effective target in `tunnelQualityProber.probeTunnel`:
|
||||
|
||||
- Type 1 and unknown tunnel fallback probes entry node to effective target instead of hardcoded Bing.
|
||||
- Type 2 probes the selected/current exit node to effective target instead of hardcoded Bing.
|
||||
- `probeBestExitOwners` receives the effective target and passes it into best-exit owner scoring.
|
||||
|
||||
Use the effective target in `evaluateBestExitOwner`:
|
||||
|
||||
- Owner-to-exit measurement stays unchanged.
|
||||
- Exit-to-public measurement probes `target.Host:target.Port` instead of `bestExitPublicTargetHost:bestExitPublicTargetPort`.
|
||||
- The per-round public probe cache key must include node ID plus target host and port so future extensions cannot reuse measurements across different targets.
|
||||
|
||||
## API Shape
|
||||
|
||||
Tunnel list/get data includes:
|
||||
|
||||
```json
|
||||
{
|
||||
"probeTargetHost": "example.com",
|
||||
"probeTargetPort": 443
|
||||
}
|
||||
```
|
||||
|
||||
For old/default tunnels, return empty host and `0` to represent `use default`. The edit form must preserve default-as-empty unless the user explicitly saves a custom target.
|
||||
|
||||
Quality monitoring response includes effective target display metadata:
|
||||
|
||||
```json
|
||||
{
|
||||
"probeTargetHost": "www.bing.com",
|
||||
"probeTargetPort": 443
|
||||
}
|
||||
```
|
||||
|
||||
Existing `exitToBingLatency` and `exitToBingLoss` keys stay to avoid breaking frontend and external consumers.
|
||||
|
||||
## Frontend Flow
|
||||
|
||||
Extend `ChainTunnel` only if needed for node-level data; the target belongs to the tunnel, so `Tunnel` and `TunnelForm` get:
|
||||
|
||||
- `probeTargetHost?: string`
|
||||
- `probeTargetPort?: number`
|
||||
|
||||
On edit:
|
||||
|
||||
- Populate form fields from tunnel response.
|
||||
- Empty or zero means default target.
|
||||
|
||||
On submit:
|
||||
|
||||
- Trim host.
|
||||
- Convert blank port to `0`.
|
||||
- Send `probeTargetHost` and `probeTargetPort` with create/update payload.
|
||||
|
||||
Display:
|
||||
|
||||
- In the form helper, show default target behavior.
|
||||
- In quality/best-exit display areas, avoid `Bing` wording; prefer `测试目标` or the concrete `host:port`.
|
||||
|
||||
## Error Handling
|
||||
|
||||
- Invalid target input returns a normal API error envelope with a specific message.
|
||||
- Probe failures use existing quality error paths and best-exit scoring failure entries.
|
||||
- If all exit-to-target probes fail, best-exit behavior remains the same as today when all Bing probes fail: no valid best decision is applied from that round.
|
||||
|
||||
## Testing
|
||||
|
||||
Backend tests:
|
||||
|
||||
- Normalize default target when host/port are empty.
|
||||
- Reject partial host/port configuration and invalid port ranges.
|
||||
- Reject host values with URL scheme/path/whitespace.
|
||||
- Create/update tunnel persists `probeTargetHost` and `probeTargetPort`.
|
||||
- `ListTunnels` returns target fields.
|
||||
- `tunnelQualityProber` uses configured target instead of `www.bing.com:443`.
|
||||
- `best` scoring uses configured target for exit-to-target probes.
|
||||
- Empty target preserves old default `www.bing.com:443` behavior.
|
||||
|
||||
Frontend verification:
|
||||
|
||||
- `pnpm run build` passes.
|
||||
- Manual UI check: create/edit tunnel with blank target and custom target, confirm payload and round-trip display.
|
||||
|
||||
## Rollout And Compatibility
|
||||
|
||||
- Existing tunnels continue using `www.bing.com:443` because empty target resolves to default.
|
||||
- SQLite/PostgreSQL schema changes are handled by existing auto-migration.
|
||||
- Historical `TunnelQuality` rows keep existing columns and are not rewritten.
|
||||
- No runtime agent change is required; the panel already performs these quality probes through existing node ping APIs.
|
||||
|
||||
## Open Decisions
|
||||
|
||||
None. User-approved decisions:
|
||||
|
||||
- Per-tunnel fields are `host + port`.
|
||||
- The target applies to both `best` scoring and tunnel quality monitoring.
|
||||
- Probe type remains TCP host/port.
|
||||
- Empty target defaults to `www.bing.com:443`.
|
||||
@@ -9,10 +9,14 @@ ARG TARGETOS
|
||||
ARG TARGETARCH
|
||||
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -o /out/paneld ./cmd/paneld
|
||||
|
||||
FROM docker:27-cli AS dockercli
|
||||
|
||||
FROM debian:bookworm-slim
|
||||
WORKDIR /app
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates wget && rm -rf /var/lib/apt/lists/*
|
||||
COPY --from=builder /out/paneld /app/paneld
|
||||
COPY --from=dockercli /usr/local/bin/docker /usr/local/bin/docker
|
||||
COPY --from=dockercli /usr/local/libexec/docker/cli-plugins/docker-compose /usr/local/libexec/docker/cli-plugins/docker-compose
|
||||
|
||||
ENV SERVER_ADDR=:6365
|
||||
EXPOSE 6365
|
||||
|
||||
@@ -90,7 +90,7 @@
|
||||
| 表名 | Model | 特殊处理 |
|
||||
|------|-------|----------|
|
||||
| `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) |
|
||||
| `forward` | `Forward` | |
|
||||
| `forward` | `Forward` | 增加 `proxy_protocol` 字段 |
|
||||
| `forward_port` | `ForwardPort` | |
|
||||
| `node` | `Node` | |
|
||||
| `speed_limit` | `SpeedLimit` | |
|
||||
|
||||
@@ -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
|
||||
@@ -264,6 +274,13 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
speed = utSpeed
|
||||
}
|
||||
|
||||
var ipSpeed *int
|
||||
if forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0 {
|
||||
if speedVal, err := h.repo.GetSpeedLimitSpeed(forward.IPSpeedID.Int64); err == nil && speedVal > 0 {
|
||||
ipSpeed = &speedVal
|
||||
}
|
||||
}
|
||||
|
||||
serviceBase := buildForwardServiceBaseWithResolvedUserTunnel(forward.ID, forward.UserID, userTunnelID)
|
||||
|
||||
user, err := h.repo.GetUserByID(forward.UserID)
|
||||
@@ -271,19 +288,17 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var cLimiterName string
|
||||
var maxConnToSet int
|
||||
|
||||
if forward.MaxConn > 0 {
|
||||
maxConnToSet = forward.MaxConn
|
||||
cLimiterName = fmt.Sprintf("rule_conn_limit_%d", forward.ID)
|
||||
} else if user != nil && user.MaxConn > 0 {
|
||||
maxConnToSet = user.MaxConn
|
||||
cLimiterName = fmt.Sprintf("user_conn_limit_%d", user.ID)
|
||||
userMaxConn := 0
|
||||
if user != nil && user.MaxConn > 0 {
|
||||
userMaxConn = user.MaxConn
|
||||
}
|
||||
connLimiterConfigs := buildConnLimiterConfigs(forward, userMaxConn)
|
||||
|
||||
for _, fp := range ports {
|
||||
runtimeLimiters := forwardRuntimeLimiters{ConnLimiter: joinLimiterNames(connLimiterConfigs)}
|
||||
trafficLimiterNames := make([]string, 0, 2)
|
||||
if limiterID != nil && speed != nil {
|
||||
totalLimiterName := strconv.FormatInt(*limiterID, 10)
|
||||
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
||||
// If the limiter push fails because the node is offline, skip it with a warning
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
@@ -297,10 +312,29 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
trafficLimiterNames = append(trafficLimiterNames, totalLimiterName)
|
||||
}
|
||||
if ipSpeed != nil {
|
||||
ruleLimiterName := fmt.Sprintf("rule_traffic_limit_%d", forward.ID)
|
||||
if err := h.ensureTrafficLimiterOnNode(fp.NodeID, ruleLimiterName, nil, ipSpeed); err != nil {
|
||||
// If the limiter push fails because the node is offline, skip it with a warning
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
node, _ := h.getNodeRecord(fp.NodeID)
|
||||
nodeName := fmt.Sprintf("%d", fp.NodeID)
|
||||
if node != nil && strings.TrimSpace(node.Name) != "" {
|
||||
nodeName = strings.TrimSpace(node.Name)
|
||||
}
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %s 不在线,已跳过下发", nodeName))
|
||||
continue
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
trafficLimiterNames = append(trafficLimiterNames, ruleLimiterName)
|
||||
}
|
||||
runtimeLimiters.TrafficLimiter = strings.Join(trafficLimiterNames, ",")
|
||||
|
||||
if cLimiterName != "" {
|
||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, cLimiterName, maxConnToSet); err != nil {
|
||||
for _, connLimiterConfig := range connLimiterConfigs {
|
||||
if err := h.ensureConnLimiterOnNode(fp.NodeID, connLimiterConfig); err != nil {
|
||||
warnings = append(warnings, fmt.Sprintf("节点 %d 连接限制器下发失败: %v", fp.NodeID, err))
|
||||
}
|
||||
}
|
||||
@@ -309,7 +343,7 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, cLimiterName)
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), runtimeLimiters)
|
||||
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
||||
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||
if isNotFoundError(err) {
|
||||
@@ -324,7 +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, cLimiterName)
|
||||
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, runtimeLimiters)
|
||||
if err == nil && warning != "" {
|
||||
warnings = append(warnings, warning)
|
||||
}
|
||||
@@ -350,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, cLimiterName string) (string, error) {
|
||||
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, runtimeLimiters forwardRuntimeLimiters) (string, error) {
|
||||
if h == nil || forward == nil || tunnel == nil || node == nil {
|
||||
return "", errors.New("invalid bind fallback context")
|
||||
}
|
||||
@@ -367,7 +401,7 @@ func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunne
|
||||
}
|
||||
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, cLimiterName)
|
||||
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", runtimeLimiters)
|
||||
if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil {
|
||||
return "", err
|
||||
}
|
||||
@@ -450,13 +484,39 @@ func (h *Handler) forwardServiceBaseCandidates(forward *forwardRecord) ([]string
|
||||
}
|
||||
|
||||
func (h *Handler) deleteForwardServiceBasesOnNode(nodeID int64, bases []string) error {
|
||||
return deleteForwardServiceCandidates(bases, func(name string) error {
|
||||
payload := map[string]interface{}{
|
||||
"services": []string{name},
|
||||
names := buildForwardServiceDeleteNames(bases)
|
||||
if len(names) == 0 {
|
||||
return nil
|
||||
}
|
||||
payload := map[string]interface{}{"services": names}
|
||||
_, err := h.sendNodeCommand(nodeID, "DeleteService", payload, false, true)
|
||||
return err
|
||||
}
|
||||
|
||||
func buildForwardServiceDeleteNames(bases []string) []string {
|
||||
names := make([]string, 0, len(bases)*3)
|
||||
seen := make(map[string]struct{}, len(bases)*3)
|
||||
appendName := func(name string) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return
|
||||
}
|
||||
_, err := h.sendNodeCommand(nodeID, "DeleteService", payload, false, false)
|
||||
return err
|
||||
})
|
||||
if _, ok := seen[name]; ok {
|
||||
return
|
||||
}
|
||||
seen[name] = struct{}{}
|
||||
names = append(names, name)
|
||||
}
|
||||
for _, base := range bases {
|
||||
base = strings.TrimSpace(base)
|
||||
if base == "" {
|
||||
continue
|
||||
}
|
||||
appendName(base + "_tcp")
|
||||
appendName(base + "_udp")
|
||||
appendName(base)
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func (h *Handler) controlForwardServices(forward *forwardRecord, commandType string, tolerateNotFound bool) error {
|
||||
@@ -486,6 +546,25 @@ func (h *Handler) controlForwardServices(forward *forwardRecord, commandType str
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, userTunnelIDs...)
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, allUserTunnelIDs...)
|
||||
bases := buildForwardServiceBaseCandidates(forward.ID, forward.UserID, userTunnelID, candidateTunnelIDs)
|
||||
if strings.EqualFold(strings.TrimSpace(commandType), "DeleteService") {
|
||||
seen := map[int64]struct{}{}
|
||||
for _, fp := range ports {
|
||||
if _, ok := seen[fp.NodeID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[fp.NodeID] = struct{}{}
|
||||
if err := h.deleteForwardServiceBasesOnNode(fp.NodeID, bases); err != nil {
|
||||
if isNodeOfflineOrTimeoutError(err) {
|
||||
continue
|
||||
}
|
||||
if tolerateNotFound && isNotFoundError(err) {
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
seen := map[int64]struct{}{}
|
||||
healed := false
|
||||
for _, fp := range ports {
|
||||
@@ -899,6 +978,7 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
|
||||
ipPreference := h.repo.GetTunnelIPPreference(tunnelID)
|
||||
protocol := strings.ToLower(strings.TrimSpace(tunnel.Protocol))
|
||||
probeTarget := effectiveTunnelProbeTargetValues(tunnel.ProbeTargetHost, tunnel.ProbeTargetPort)
|
||||
inNodes, chainHops, outNodes := splitChainNodeGroups(chainRows)
|
||||
workItems := make([]diagnosisWorkItem, 0, len(chainRows)*2)
|
||||
|
||||
@@ -908,8 +988,8 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
|
||||
workItems = append(workItems, diagnosisWorkItem{
|
||||
fromNodeID: inNode.NodeID,
|
||||
targetIP: "www.bing.com",
|
||||
targetPort: 443,
|
||||
targetIP: probeTarget.Host,
|
||||
targetPort: probeTarget.Port,
|
||||
description: description,
|
||||
protocol: "tcp",
|
||||
metadata: map[string]interface{}{
|
||||
@@ -1000,8 +1080,8 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
description := fmt.Sprintf("出口(%s)->外网", outNode.NodeName)
|
||||
workItems = append(workItems, diagnosisWorkItem{
|
||||
fromNodeID: outNode.NodeID,
|
||||
targetIP: "www.bing.com",
|
||||
targetPort: 443,
|
||||
targetIP: probeTarget.Host,
|
||||
targetPort: probeTarget.Port,
|
||||
description: description,
|
||||
protocol: "tcp",
|
||||
metadata: map[string]interface{}{
|
||||
@@ -1014,8 +1094,8 @@ func (h *Handler) prepareTunnelDiagnosis(tunnelID int64) (string, string, []diag
|
||||
description := fmt.Sprintf("入口(%s)->外网", inNode.NodeName)
|
||||
workItems = append(workItems, diagnosisWorkItem{
|
||||
fromNodeID: inNode.NodeID,
|
||||
targetIP: "www.bing.com",
|
||||
targetPort: 443,
|
||||
targetIP: probeTarget.Host,
|
||||
targetPort: probeTarget.Port,
|
||||
description: description,
|
||||
protocol: "tcp",
|
||||
metadata: map[string]interface{}{
|
||||
@@ -1659,7 +1739,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, cLimiterName string) []map[string]interface{} {
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, runtimeLimiters forwardRuntimeLimiters) []map[string]interface{} {
|
||||
protocols := []string{"tcp", "udp"}
|
||||
services := make([]map[string]interface{}, 0, 2)
|
||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||
@@ -1702,8 +1782,18 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
},
|
||||
},
|
||||
}
|
||||
if cLimiterName != "" {
|
||||
service["climiter"] = cLimiterName
|
||||
if runtimeLimiters.ConnLimiter != "" {
|
||||
service["climiter"] = runtimeLimiters.ConnLimiter
|
||||
}
|
||||
if runtimeLimiters.TrafficLimiter != "" {
|
||||
service["limiter"] = runtimeLimiters.TrafficLimiter
|
||||
}
|
||||
if 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{}{
|
||||
@@ -1716,10 +1806,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)
|
||||
}
|
||||
@@ -1820,22 +1910,16 @@ func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxConn int) error {
|
||||
limitStr := fmt.Sprintf("$ %d", maxConn)
|
||||
|
||||
payload := map[string]interface{}{
|
||||
"name": limiterName,
|
||||
"limits": []string{limitStr},
|
||||
func (h *Handler) ensureConnLimiterOnNode(nodeID int64, cfg forwardLimiterConfig) error {
|
||||
if cfg.Name == "" || len(cfg.Limits) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
payload := map[string]interface{}{"name": cfg.Name, "limits": cfg.Limits}
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddCLimiters", payload, false, false); err != nil {
|
||||
if !isAlreadyExistsMessage(err.Error()) {
|
||||
return fmt.Errorf("连接限制器下发失败: %w", err)
|
||||
}
|
||||
updatePayload := map[string]interface{}{
|
||||
"limiter": limiterName,
|
||||
"data": payload,
|
||||
}
|
||||
updatePayload := map[string]interface{}{"limiter": cfg.Name, "data": payload}
|
||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateCLimiters", updatePayload, false, false); updateErr != nil {
|
||||
return fmt.Errorf("连接限制器更新失败: %w", updateErr)
|
||||
}
|
||||
@@ -1843,15 +1927,59 @@ func (h *Handler) ensureConnLimiterOnNode(nodeID int64, limiterName string, maxC
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildConnLimiterConfigs(forward *forwardRecord, userMaxConn int) []forwardLimiterConfig {
|
||||
if forward == nil {
|
||||
return nil
|
||||
}
|
||||
if forward.MaxConn > 0 {
|
||||
limits := []string{fmt.Sprintf("$ %d", forward.MaxConn)}
|
||||
if forward.IPMaxConn > 0 {
|
||||
limits = append(limits, fmt.Sprintf("$$ %d", forward.IPMaxConn))
|
||||
}
|
||||
return []forwardLimiterConfig{{Name: fmt.Sprintf("rule_conn_limit_%d", forward.ID), Limits: limits}}
|
||||
}
|
||||
configs := make([]forwardLimiterConfig, 0, 2)
|
||||
if userMaxConn > 0 {
|
||||
configs = append(configs, forwardLimiterConfig{Name: fmt.Sprintf("user_conn_limit_%d", forward.UserID), Limits: []string{fmt.Sprintf("$ %d", userMaxConn)}})
|
||||
}
|
||||
if forward.IPMaxConn > 0 {
|
||||
configs = append(configs, forwardLimiterConfig{Name: fmt.Sprintf("rule_conn_limit_%d", forward.ID), Limits: []string{fmt.Sprintf("$$ %d", forward.IPMaxConn)}})
|
||||
}
|
||||
return configs
|
||||
}
|
||||
|
||||
func joinLimiterNames(configs []forwardLimiterConfig) string {
|
||||
names := make([]string, 0, len(configs))
|
||||
for _, cfg := range configs {
|
||||
if cfg.Name != "" {
|
||||
names = append(names, cfg.Name)
|
||||
}
|
||||
}
|
||||
return strings.Join(names, ",")
|
||||
}
|
||||
|
||||
func speedToLimitLine(key string, speed int) string {
|
||||
rate := float64(speed) / 8.0
|
||||
return fmt.Sprintf("%s %.1fMB %.1fMB", key, rate, rate)
|
||||
}
|
||||
|
||||
func buildTrafficLimiterPayload(name string, totalSpeed *int, ipSpeed *int) map[string]interface{} {
|
||||
limits := make([]string, 0, 3)
|
||||
if totalSpeed != nil && *totalSpeed > 0 {
|
||||
limits = append(limits, speedToLimitLine("$", *totalSpeed))
|
||||
}
|
||||
if ipSpeed != nil && *ipSpeed > 0 {
|
||||
limits = append(limits, speedToLimitLine("0.0.0.0/0", *ipSpeed), speedToLimitLine("::/0", *ipSpeed))
|
||||
}
|
||||
return map[string]interface{}{"name": name, "limits": limits}
|
||||
}
|
||||
|
||||
func buildLimiterAddPayload(limiterID int64, speed int) (string, map[string]interface{}) {
|
||||
rate := float64(speed) / 8.0
|
||||
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
|
||||
name := strconv.FormatInt(limiterID, 10)
|
||||
|
||||
return name, map[string]interface{}{
|
||||
"name": name,
|
||||
"limits": []string{limitStr},
|
||||
"limits": []string{speedToLimitLine("$", speed)},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1879,3 +2007,20 @@ func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureTrafficLimiterOnNode(nodeID int64, name string, totalSpeed *int, ipSpeed *int) error {
|
||||
payload := buildTrafficLimiterPayload(name, totalSpeed, ipSpeed)
|
||||
limits, _ := payload["limits"].([]string)
|
||||
if name == "" || len(limits) == 0 {
|
||||
return nil
|
||||
}
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
|
||||
if !isAlreadyExistsMessage(err.Error()) {
|
||||
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||
}
|
||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil {
|
||||
return fmt.Errorf("限速规则更新失败: %w", updateErr)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -201,6 +201,52 @@ func TestDeleteForwardServiceCandidatesDeletesAllMatchingVariants(t *testing.T)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceDeleteNamesBatchesAndDeduplicatesVariants(t *testing.T) {
|
||||
bases := []string{"57_7_7", "57_7_0", "57_7_7"}
|
||||
got := buildForwardServiceDeleteNames(bases)
|
||||
want := []string{"57_7_7_tcp", "57_7_7_udp", "57_7_7", "57_7_0_tcp", "57_7_0_udp", "57_7_0"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("expected %v, got %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemovedTunnelRuntimeNodeIDsSeparatesChainAndServiceRoles(t *testing.T) {
|
||||
oldRows := []chainNodeRecord{
|
||||
{NodeID: 1, ChainType: 1},
|
||||
{NodeID: 2, ChainType: 2},
|
||||
{NodeID: 3, ChainType: 3},
|
||||
{NodeID: 5, ChainType: 2},
|
||||
{NodeID: 6, ChainType: 3},
|
||||
}
|
||||
newRows := []chainNodeRecord{
|
||||
{NodeID: 2, ChainType: 3},
|
||||
{NodeID: 3, ChainType: 3},
|
||||
{NodeID: 5, ChainType: 1},
|
||||
}
|
||||
|
||||
removedChains := removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsChain)
|
||||
if want := []int64{1, 2}; !reflect.DeepEqual(removedChains, want) {
|
||||
t.Fatalf("expected removed chains %v, got %v", want, removedChains)
|
||||
}
|
||||
|
||||
removedServices := removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsService)
|
||||
if want := []int64{5, 6}; !reflect.DeepEqual(removedServices, want) {
|
||||
t.Fatalf("expected removed services %v, got %v", want, removedServices)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelForwardRuntimeNeedsSyncOnlyWhenTypeOrEntriesChange(t *testing.T) {
|
||||
if tunnelForwardRuntimeNeedsSync(2, 2, []int64{1, 2}, []int64{2, 1}) {
|
||||
t.Fatalf("same tunnel type and same entry set should not resync forwards")
|
||||
}
|
||||
if !tunnelForwardRuntimeNeedsSync(1, 2, []int64{1}, []int64{1}) {
|
||||
t.Fatalf("type change should resync forwards")
|
||||
}
|
||||
if !tunnelForwardRuntimeNeedsSync(2, 2, []int64{1}, []int64{1, 2}) {
|
||||
t.Fatalf("entry set change should resync forwards")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateForwardPortAvailabilityRejectsOtherForwardOccupancy(t *testing.T) {
|
||||
h := &Handler{repo: nil}
|
||||
node := &nodeRecord{ID: 9, Name: "test-node"}
|
||||
@@ -378,7 +424,7 @@ func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -393,7 +439,7 @@ func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -409,7 +455,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
|
||||
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -464,7 +510,7 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -478,6 +524,58 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildConnLimiterConfigCombinesTotalAndPerIP(t *testing.T) {
|
||||
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, MaxConn: 100, IPMaxConn: 5}, 37)
|
||||
want := []forwardLimiterConfig{{Name: "rule_conn_limit_42", Limits: []string{"$ 100", "$$ 5"}}}
|
||||
if !reflect.DeepEqual(cfgs, want) {
|
||||
t.Fatalf("expected %+v, got %+v", want, cfgs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildConnLimiterConfigUsesUserTotalWithRulePerIP(t *testing.T) {
|
||||
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, IPMaxConn: 5}, 37)
|
||||
want := []forwardLimiterConfig{
|
||||
{Name: "user_conn_limit_9", Limits: []string{"$ 37"}},
|
||||
{Name: "rule_conn_limit_42", Limits: []string{"$$ 5"}},
|
||||
}
|
||||
if !reflect.DeepEqual(cfgs, want) {
|
||||
t.Fatalf("expected %+v, got %+v", want, cfgs)
|
||||
}
|
||||
if got := joinLimiterNames(cfgs); got != "user_conn_limit_9,rule_conn_limit_42" {
|
||||
t.Fatalf("expected composite limiter names, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficLimiterPayloadUsesOnlyPerIPRulesWhenTotalIsSeparate(t *testing.T) {
|
||||
payload := buildTrafficLimiterPayload("rule_traffic_limit_42", nil, intPtr(40))
|
||||
wantLimits := []string{"0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
|
||||
if payload["name"] != "rule_traffic_limit_42" {
|
||||
t.Fatalf("expected name rule_traffic_limit_42, got %v", payload["name"])
|
||||
}
|
||||
if !reflect.DeepEqual(payload["limits"], wantLimits) {
|
||||
t.Fatalf("expected limits %v, got %v", wantLimits, payload["limits"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigsUsesRuntimeLimiterNames(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{TrafficLimiter: "rule_traffic_limit_42", ConnLimiter: "rule_conn_limit_42"})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
for _, service := range services {
|
||||
if service["limiter"] != "rule_traffic_limit_42" {
|
||||
t.Fatalf("expected traffic limiter rule_traffic_limit_42, got %v", service["limiter"])
|
||||
}
|
||||
if service["climiter"] != "rule_conn_limit_42" {
|
||||
t.Fatalf("expected conn limiter rule_conn_limit_42, got %v", service["climiter"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func intPtr(v int) *int { return &v }
|
||||
|
||||
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -85,6 +85,81 @@ type federationRuntimeReleaseRoleRequest struct {
|
||||
ResourceKey string `json:"resourceKey"`
|
||||
}
|
||||
|
||||
func federationRuntimeChainName(bindingID string) string {
|
||||
bindingID = strings.TrimSpace(bindingID)
|
||||
if bindingID == "" {
|
||||
return ""
|
||||
}
|
||||
return "fed_chain_" + bindingID
|
||||
}
|
||||
|
||||
func buildFederationMiddleChainConfig(chainName string, runtimeID int64, protocol, strategy string, targets []federationRuntimeTarget, interfaceName string) (map[string]interface{}, error) {
|
||||
chainName = strings.TrimSpace(chainName)
|
||||
if chainName == "" {
|
||||
return nil, fmt.Errorf("chain name is required")
|
||||
}
|
||||
if len(targets) == 0 {
|
||||
return nil, fmt.Errorf("targets are required for middle role")
|
||||
}
|
||||
protocol = defaultString(protocol, "tls")
|
||||
nodeItems := make([]map[string]interface{}, 0, len(targets))
|
||||
for i, target := range targets {
|
||||
host := strings.TrimSpace(target.Host)
|
||||
if host == "" || target.Port <= 0 {
|
||||
return nil, fmt.Errorf("Invalid target")
|
||||
}
|
||||
targetProtocol := defaultString(target.Protocol, protocol)
|
||||
connector := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"nodelay": true,
|
||||
"mux.keepaliveInterval": "15s",
|
||||
"mux.keepaliveTimeout": "45s",
|
||||
}
|
||||
}
|
||||
if isKCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"connectTimeout": "30s",
|
||||
}
|
||||
}
|
||||
nodeItems = append(nodeItems, map[string]interface{}{
|
||||
"name": fmt.Sprintf("node_%d", i+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
|
||||
"connector": connector,
|
||||
"dialer": buildTunnelDialerConfig(targetProtocol),
|
||||
})
|
||||
}
|
||||
|
||||
chainData := map[string]interface{}{
|
||||
"name": chainName,
|
||||
"hops": []map[string]interface{}{
|
||||
{
|
||||
"name": fmt.Sprintf("hop_%d", runtimeID),
|
||||
"selector": map[string]interface{}{
|
||||
"strategy": runtimeTunnelStrategy(strategy),
|
||||
"maxFails": 1,
|
||||
"failTimeout": int64(600000000000),
|
||||
},
|
||||
"nodes": nodeItems,
|
||||
},
|
||||
},
|
||||
}
|
||||
if strings.TrimSpace(interfaceName) != "" {
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
hops[0]["interface"] = interfaceName
|
||||
}
|
||||
return chainData, nil
|
||||
}
|
||||
|
||||
func updateChainPayload(chainName string, chainData map[string]interface{}) map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"chain": chainName,
|
||||
"data": chainData,
|
||||
}
|
||||
}
|
||||
|
||||
type federationRuntimeDiagnoseRequest struct {
|
||||
IP string `json:"ip"`
|
||||
Port int `json:"port"`
|
||||
@@ -145,8 +220,8 @@ type remoteUsageNodeItem struct {
|
||||
|
||||
func buildFederationServiceConfig(serviceName, addr, protocol, role, chainName string, targetCount int, interfaceName string) map[string]interface{} {
|
||||
service := map[string]interface{}{
|
||||
"name": serviceName,
|
||||
"addr": addr,
|
||||
"name": serviceName,
|
||||
"addr": addr,
|
||||
"handler": map[string]interface{}{
|
||||
"type": "relay",
|
||||
},
|
||||
@@ -1056,7 +1131,43 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
return
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(share.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
protocol := defaultString(req.Protocol, runtime.Protocol)
|
||||
strategy := defaultString(req.Strategy, "round")
|
||||
chainName := defaultString(runtime.ChainName, federationRuntimeChainName(runtime.BindingID))
|
||||
if chainName == "" {
|
||||
chainName = federationRuntimeChainName(fmt.Sprintf("%d", runtime.ID))
|
||||
}
|
||||
serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID)
|
||||
if runtime.Applied == 1 && strings.TrimSpace(runtime.BindingID) != "" {
|
||||
if req.Role == "middle" && len(req.Targets) > 0 {
|
||||
chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName)
|
||||
if buildErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(buildErr.Error()))
|
||||
return
|
||||
}
|
||||
if _, err := h.sendNodeCommand(share.NodeID, "UpdateChains", updateChainPayload(chainName, chainData), false, false); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
targetBytes, _ := json.Marshal(req.Targets)
|
||||
runtime.Role = req.Role
|
||||
runtime.ChainName = chainName
|
||||
runtime.Protocol = protocol
|
||||
runtime.Strategy = strategy
|
||||
runtime.Target = string(targetBytes)
|
||||
runtime.Status = 1
|
||||
runtime.UpdatedTime = time.Now().UnixMilli()
|
||||
if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"bindingId": runtime.BindingID,
|
||||
"allocatedPort": runtime.Port,
|
||||
@@ -1076,71 +1187,12 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ
|
||||
}
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(share.NodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
protocol := defaultString(req.Protocol, runtime.Protocol)
|
||||
strategy := defaultString(req.Strategy, "round")
|
||||
chainName := fmt.Sprintf("fed_chain_%d", runtime.ID)
|
||||
serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID)
|
||||
|
||||
if req.Role == "middle" {
|
||||
if len(req.Targets) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("targets are required for middle role"))
|
||||
chainData, buildErr := buildFederationMiddleChainConfig(chainName, runtime.ID, protocol, strategy, req.Targets, node.InterfaceName)
|
||||
if buildErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(buildErr.Error()))
|
||||
return
|
||||
}
|
||||
nodeItems := make([]map[string]interface{}, 0, len(req.Targets))
|
||||
for i, target := range req.Targets {
|
||||
host := strings.TrimSpace(target.Host)
|
||||
if host == "" || target.Port <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid target"))
|
||||
return
|
||||
}
|
||||
targetProtocol := defaultString(target.Protocol, protocol)
|
||||
connector := map[string]interface{}{
|
||||
"type": "relay",
|
||||
}
|
||||
if isTCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"nodelay": true,
|
||||
"mux.keepaliveInterval": "15s",
|
||||
"mux.keepaliveTimeout": "45s",
|
||||
}
|
||||
}
|
||||
if isKCPTunnelProtocol(targetProtocol) {
|
||||
connector["metadata"] = map[string]interface{}{
|
||||
"connectTimeout": "30s",
|
||||
}
|
||||
}
|
||||
nodeItems = append(nodeItems, map[string]interface{}{
|
||||
"name": fmt.Sprintf("node_%d", i+1),
|
||||
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
|
||||
"connector": connector,
|
||||
"dialer": buildTunnelDialerConfig(targetProtocol),
|
||||
})
|
||||
}
|
||||
|
||||
chainData := map[string]interface{}{
|
||||
"name": chainName,
|
||||
"hops": []map[string]interface{}{
|
||||
{
|
||||
"name": fmt.Sprintf("hop_%d", runtime.ID),
|
||||
"selector": map[string]interface{}{
|
||||
"strategy": strategy,
|
||||
"maxFails": 1,
|
||||
"failTimeout": int64(600000000000),
|
||||
},
|
||||
"nodes": nodeItems,
|
||||
},
|
||||
},
|
||||
}
|
||||
if strings.TrimSpace(node.InterfaceName) != "" {
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
hops[0]["interface"] = node.InterfaceName
|
||||
}
|
||||
if _, err := h.sendNodeCommand(share.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
|
||||
@@ -281,6 +281,63 @@ func TestBuildFederationServiceConfig_NonTLSProtocol_NoNodelay(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeChainNameDerivesFromBindingID(t *testing.T) {
|
||||
if got := federationRuntimeChainName("12"); got != "fed_chain_12" {
|
||||
t.Fatalf("expected fed_chain_12, got %q", got)
|
||||
}
|
||||
if got := federationRuntimeChainName(" 12 "); got != "fed_chain_12" {
|
||||
t.Fatalf("expected trimmed fed_chain_12, got %q", got)
|
||||
}
|
||||
if got := federationRuntimeChainName(""); got != "" {
|
||||
t.Fatalf("expected blank binding ID to stay blank, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFederationMiddleChainConfigUsesExistingChainNameAndBestStrategy(t *testing.T) {
|
||||
chainData, err := buildFederationMiddleChainConfig("fed_chain_12", 12, "tls", tunnelStrategyBest, []federationRuntimeTarget{
|
||||
{Host: "10.0.0.31", Port: 30031, Protocol: "tls"},
|
||||
{Host: "10.0.0.30", Port: 30030, Protocol: "tls"},
|
||||
}, "")
|
||||
if err != nil {
|
||||
t.Fatalf("build chain: %v", err)
|
||||
}
|
||||
if chainData["name"] != "fed_chain_12" {
|
||||
t.Fatalf("expected existing chain name, got %v", chainData["name"])
|
||||
}
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
selector := hops[0]["selector"].(map[string]interface{})
|
||||
if selector["strategy"] != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best strategy to map to fifo, got %v", selector["strategy"])
|
||||
}
|
||||
nodes := hops[0]["nodes"].([]map[string]interface{})
|
||||
if nodes[0]["addr"] != "10.0.0.31:30031" || nodes[1]["addr"] != "10.0.0.30:30030" {
|
||||
t.Fatalf("expected target order to be preserved, got %+v", nodes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateChainPayloadWrapsChainDataForAgentUpdate(t *testing.T) {
|
||||
chainData := map[string]interface{}{
|
||||
"name": "fed_chain_12",
|
||||
"hops": []map[string]interface{}{},
|
||||
}
|
||||
|
||||
payload := updateChainPayload("fed_chain_12", chainData)
|
||||
if len(payload) != 2 {
|
||||
t.Fatalf("expected exact wrapper with 2 keys, got %+v", payload)
|
||||
}
|
||||
if payload["chain"] != "fed_chain_12" {
|
||||
t.Fatalf("expected chain name in wrapper, got %v", payload["chain"])
|
||||
}
|
||||
wrappedData, ok := payload["data"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected wrapped chain data map, got %T", payload["data"])
|
||||
}
|
||||
chainData["name"] = "fed_chain_12_updated"
|
||||
if wrappedData["name"] != "fed_chain_12_updated" {
|
||||
t.Fatalf("expected wrapper to preserve chainData identity, got %+v", wrappedData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFederationRuntimeReservePortRejectsWhenShareFlowExceeded(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sort"
|
||||
@@ -21,6 +22,7 @@ import (
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
@@ -43,13 +45,19 @@ type Handler struct {
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
|
||||
upgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
upgradeMu sync.Mutex
|
||||
systemUpgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
nodeOnlineRedeployAt map[int64]time.Time
|
||||
nodeOnlineRedeployQueued map[int64]struct{}
|
||||
nodeOnlineRedeploying map[int64]struct{}
|
||||
|
||||
qualityProber *tunnelQualityProber
|
||||
bestExit *bestExitManager
|
||||
}
|
||||
|
||||
const monitorTunnelQualityEnabledConfigKey = "monitor_tunnel_quality_enabled"
|
||||
const allowLocalRemoteAddrConfigKey = "allow_local_remote_addr"
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
@@ -95,13 +103,17 @@ 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{}),
|
||||
bestExit: newBestExitManager(),
|
||||
}
|
||||
h.healthCheck = health.NewChecker(repo, h.wsServer)
|
||||
h.qualityProber = newTunnelQualityProber(h)
|
||||
@@ -144,6 +156,10 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
|
||||
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
|
||||
mux.HandleFunc("/api/v1/system/storage", h.storageSummary)
|
||||
mux.HandleFunc("/api/v1/system/version", h.systemVersion)
|
||||
mux.HandleFunc("/api/v1/system/check-updates", h.systemCheckUpdates)
|
||||
mux.HandleFunc("/api/v1/system/upgrade", h.systemUpgrade)
|
||||
mux.HandleFunc("/api/v1/license/activate", h.licenseActivate)
|
||||
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
|
||||
@@ -463,6 +479,7 @@ func (h *Handler) tunnelList(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStates(items)
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
@@ -796,11 +813,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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1023,6 +1045,8 @@ func normalizeAndValidateConfigValue(key, value string) (string, error) {
|
||||
default:
|
||||
return "", fmt.Errorf("隧道质量检测开关配置值无效")
|
||||
}
|
||||
case monitoring.ConfigMonitorRetentionDays:
|
||||
return monitoring.NormalizeMonitoringRetentionDays(value)
|
||||
default:
|
||||
return value, nil
|
||||
}
|
||||
@@ -1041,6 +1065,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("请求失败"))
|
||||
|
||||
@@ -234,8 +234,22 @@ func (h *Handler) monitorTunnelQualityHandler(w http.ResponseWriter, r *http.Req
|
||||
return
|
||||
}
|
||||
|
||||
targetsByTunnelID := map[int64]tunnelProbeTarget{}
|
||||
if tunnels, listErr := h.repo.ListTunnels(); listErr == nil {
|
||||
for _, item := range tunnels {
|
||||
id := asInt64(item["id"], 0)
|
||||
if id > 0 {
|
||||
targetsByTunnelID[id] = effectiveTunnelProbeTargetValues(asString(item["probeTargetHost"]), asInt(item["probeTargetPort"], 0))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
snapshots := make([]tunnelQualitySnapshot, 0, len(qualities))
|
||||
for _, q := range qualities {
|
||||
target := targetsByTunnelID[q.TunnelID]
|
||||
if target.Host == "" {
|
||||
target = defaultTunnelProbeTarget()
|
||||
}
|
||||
snapshots = append(snapshots, tunnelQualitySnapshot{
|
||||
TunnelID: q.TunnelID,
|
||||
EntryToExitLatency: q.EntryToExitLatency,
|
||||
@@ -246,6 +260,8 @@ func (h *Handler) monitorTunnelQualityHandler(w http.ResponseWriter, r *http.Req
|
||||
ErrorMessage: q.ErrorMessage,
|
||||
Timestamp: q.Timestamp,
|
||||
ChainDetails: q.ChainDetails,
|
||||
ProbeTargetHost: target.Host,
|
||||
ProbeTargetPort: target.Port,
|
||||
})
|
||||
}
|
||||
response.WriteJSON(w, response.OK(snapshots))
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"math/big"
|
||||
"net"
|
||||
"net/http"
|
||||
@@ -138,6 +139,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 {
|
||||
@@ -210,6 +220,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())
|
||||
}
|
||||
@@ -586,6 +607,17 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
trafficRatio := asFloat(req["trafficRatio"], 1.0)
|
||||
inIP := asString(req["inIp"])
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
probeTarget, probeTargetConfigured, err := parseTunnelProbeTargetFromRequest(req)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
probeTargetHost := ""
|
||||
probeTargetPort := 0
|
||||
if probeTargetConfigured {
|
||||
probeTargetHost = probeTarget.Host
|
||||
probeTargetPort = probeTarget.Port
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
inx := h.repo.NextIndex("tunnel")
|
||||
localDomain := h.federationLocalDomain()
|
||||
@@ -664,17 +696,19 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
tunnelProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
|
||||
}
|
||||
tunnel := model.Tunnel{
|
||||
Name: name,
|
||||
TrafficRatio: trafficRatio,
|
||||
Type: typeVal,
|
||||
Protocol: tunnelProtocol,
|
||||
Flow: flow,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: status,
|
||||
InIP: tunnelInIP,
|
||||
Inx: inx,
|
||||
IPPreference: ipPreference,
|
||||
Name: name,
|
||||
TrafficRatio: trafficRatio,
|
||||
Type: typeVal,
|
||||
Protocol: tunnelProtocol,
|
||||
Flow: flow,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: status,
|
||||
InIP: tunnelInIP,
|
||||
Inx: inx,
|
||||
IPPreference: ipPreference,
|
||||
ProbeTargetHost: probeTargetHost,
|
||||
ProbeTargetPort: probeTargetPort,
|
||||
}
|
||||
if err := tx.Create(&tunnel).Error; err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -728,21 +762,8 @@ func (h *Handler) cleanupTunnelRuntime(tunnelID int64) {
|
||||
return
|
||||
}
|
||||
|
||||
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),
|
||||
}
|
||||
serviceNames := tunnelRuntimeServiceNames(tunnelID)
|
||||
|
||||
for _, row := range chainRows {
|
||||
if row.ChainType == 1 {
|
||||
@@ -756,6 +777,70 @@ func (h *Handler) cleanupTunnelRuntime(tunnelID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
func tunnelRuntimeServiceNames(tunnelID int64) []string {
|
||||
return []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),
|
||||
}
|
||||
}
|
||||
|
||||
func tunnelRuntimeNeedsChain(row chainNodeRecord) bool {
|
||||
return row.ChainType == 1 || row.ChainType == 2
|
||||
}
|
||||
|
||||
func tunnelRuntimeNeedsService(row chainNodeRecord) bool {
|
||||
return row.ChainType == 2 || row.ChainType == 3
|
||||
}
|
||||
|
||||
func removedTunnelRuntimeNodeIDs(oldRows, newRows []chainNodeRecord, needsRuntime func(chainNodeRecord) bool) []int64 {
|
||||
if len(oldRows) == 0 || needsRuntime == nil {
|
||||
return nil
|
||||
}
|
||||
newRuntimeNodes := make(map[int64]struct{}, len(newRows))
|
||||
for _, row := range newRows {
|
||||
if row.NodeID <= 0 || !needsRuntime(row) {
|
||||
continue
|
||||
}
|
||||
newRuntimeNodes[row.NodeID] = struct{}{}
|
||||
}
|
||||
seen := make(map[int64]struct{}, len(oldRows))
|
||||
removed := make([]int64, 0)
|
||||
for _, row := range oldRows {
|
||||
if row.NodeID <= 0 || !needsRuntime(row) {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[row.NodeID]; ok {
|
||||
continue
|
||||
}
|
||||
seen[row.NodeID] = struct{}{}
|
||||
if _, stillNeeded := newRuntimeNodes[row.NodeID]; stillNeeded {
|
||||
continue
|
||||
}
|
||||
removed = append(removed, row.NodeID)
|
||||
}
|
||||
return removed
|
||||
}
|
||||
|
||||
func (h *Handler) cleanupObsoleteTunnelRuntime(tunnelID int64, oldRows, newRows []chainNodeRecord) {
|
||||
if h == nil || tunnelID <= 0 || len(oldRows) == 0 {
|
||||
return
|
||||
}
|
||||
chainName := fmt.Sprintf("chains_%d", tunnelID)
|
||||
for _, nodeID := range removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsChain) {
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
|
||||
}
|
||||
serviceNames := tunnelRuntimeServiceNames(tunnelID)
|
||||
for _, nodeID := range removedTunnelRuntimeNodeIDs(oldRows, newRows, tunnelRuntimeNeedsService) {
|
||||
_, _ = h.sendNodeCommand(nodeID, "DeleteService", map[string]interface{}{"services": serviceNames}, false, true)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelGet(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -770,6 +855,7 @@ func (h *Handler) tunnelGet(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.attachBestExitStates(items)
|
||||
for _, it := range items {
|
||||
if asInt64(it["id"], 0) == id {
|
||||
response.WriteJSON(w, response.OK(it))
|
||||
@@ -798,14 +884,37 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
|
||||
return
|
||||
}
|
||||
typeVal := asInt(req["type"], 1)
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
_, hasProbeTargetHost := req["probeTargetHost"]
|
||||
_, hasProbeTargetPort := req["probeTargetPort"]
|
||||
probeTargetFieldsPresent := hasProbeTargetHost || hasProbeTargetPort
|
||||
probeTargetHost := ""
|
||||
probeTargetPort := 0
|
||||
if probeTargetFieldsPresent {
|
||||
probeTarget, probeTargetConfigured, err := parseTunnelProbeTargetFromRequest(req)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if probeTargetConfigured {
|
||||
probeTargetHost = probeTarget.Host
|
||||
probeTargetPort = probeTarget.Port
|
||||
}
|
||||
}
|
||||
oldEntryNodeIDs, _ := h.tunnelEntryNodeIDs(id)
|
||||
|
||||
h.cleanupTunnelRuntime(id)
|
||||
oldTunnel, _ := h.getTunnelRecord(id)
|
||||
if !probeTargetFieldsPresent && oldTunnel != nil {
|
||||
probeTargetHost = oldTunnel.ProbeTargetHost
|
||||
probeTargetPort = oldTunnel.ProbeTargetPort
|
||||
}
|
||||
oldChainRows, _ := h.listChainNodesForTunnel(id)
|
||||
if oldTunnel != nil && oldTunnel.Type == 2 && typeVal != 2 {
|
||||
h.cleanupTunnelRuntime(id)
|
||||
}
|
||||
h.cleanupFederationRuntime(id)
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
typeVal := asInt(req["type"], 1)
|
||||
ipPreference := asString(req["ipPreference"])
|
||||
localDomain := h.federationLocalDomain()
|
||||
|
||||
runtimeState, err := h.prepareTunnelCreateState(h.repo.DB(), req, typeVal, id)
|
||||
@@ -852,6 +961,8 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
inIp,
|
||||
ipPreference,
|
||||
updateProtocol,
|
||||
probeTargetHost,
|
||||
probeTargetPort,
|
||||
now,
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -897,13 +1008,19 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if typeVal == 2 {
|
||||
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
|
||||
applyRuntime := h.applyTunnelRuntime
|
||||
if oldTunnel != nil && oldTunnel.Type == 2 {
|
||||
applyRuntime = h.applyTunnelRuntimeUpsert
|
||||
}
|
||||
createdChains, createdServices, applyErr := applyRuntime(runtimeState)
|
||||
if applyErr != nil {
|
||||
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)
|
||||
if oldTunnel == nil || oldTunnel.Type != 2 {
|
||||
h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol)
|
||||
}
|
||||
h.releaseFederationRuntimeRefs(federationReleaseRefs)
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(id)
|
||||
if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) {
|
||||
@@ -913,9 +1030,20 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
|
||||
return
|
||||
}
|
||||
newChainRows, _ := h.listChainNodesForTunnel(id)
|
||||
h.cleanupObsoleteTunnelRuntime(id, oldChainRows, newChainRows)
|
||||
}
|
||||
|
||||
if forwards, fwdErr := h.listForwardsByTunnel(id); fwdErr == nil {
|
||||
oldType := 0
|
||||
if oldTunnel != nil {
|
||||
oldType = oldTunnel.Type
|
||||
}
|
||||
if tunnelForwardRuntimeNeedsSync(oldType, typeVal, oldEntryNodeIDs, newEntryNodeIDs) {
|
||||
forwards, fwdErr := h.listForwardsByTunnel(id)
|
||||
if fwdErr != nil {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
return
|
||||
}
|
||||
for i := range forwards {
|
||||
_ = h.syncForwardServices(&forwards[i], "UpdateService", true)
|
||||
}
|
||||
@@ -924,6 +1052,13 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func tunnelForwardRuntimeNeedsSync(oldType, newType int, oldEntryNodeIDs, newEntryNodeIDs []int64) bool {
|
||||
if oldType != newType {
|
||||
return true
|
||||
}
|
||||
return !sameInt64Set(oldEntryNodeIDs, newEntryNodeIDs)
|
||||
}
|
||||
|
||||
func sameInt64Set(a, b []int64) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
@@ -1726,11 +1861,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
|
||||
@@ -1742,6 +1879,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)
|
||||
@@ -1780,7 +1929,13 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
userName = "user"
|
||||
}
|
||||
maxConn := asInt(req["maxConn"], 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 := 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
|
||||
@@ -1851,9 +2006,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
|
||||
}
|
||||
}
|
||||
@@ -1879,6 +2034,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 {
|
||||
@@ -1932,8 +2108,13 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
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); err != nil {
|
||||
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
|
||||
}
|
||||
@@ -3170,6 +3351,8 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
nextTargets := state.OutNodes
|
||||
if hopIdx+1 < len(state.ChainHops) {
|
||||
nextTargets = state.ChainHops[hopIdx+1]
|
||||
} else {
|
||||
nextTargets = h.orderBestExitTargets(state.TunnelID, chainNode.NodeID, nextTargets)
|
||||
}
|
||||
applyTargets := make([]client.RuntimeTarget, 0, len(nextTargets))
|
||||
for _, target := range nextTargets {
|
||||
@@ -3199,7 +3382,7 @@ func (h *Handler) applyFederationRuntime(state *tunnelCreateState, localDomain s
|
||||
ResourceKey: resourceKey,
|
||||
Role: "middle",
|
||||
Protocol: defaultString(chainNode.Protocol, "tls"),
|
||||
Strategy: defaultString(chainNode.Strategy, "round"),
|
||||
Strategy: runtimeStrategyForTargets(chainNode, nextTargets),
|
||||
Targets: applyTargets,
|
||||
}
|
||||
applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq)
|
||||
@@ -3292,6 +3475,14 @@ func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) {
|
||||
return h.applyTunnelRuntimeWithMode(state, false)
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntimeUpsert(state *tunnelCreateState) ([]int64, []int64, error) {
|
||||
return h.applyTunnelRuntimeWithMode(state, true)
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelRuntimeWithMode(state *tunnelCreateState, upsert bool) ([]int64, []int64, error) {
|
||||
if h == nil || state == nil {
|
||||
return nil, nil, errors.New("invalid tunnel runtime state")
|
||||
}
|
||||
@@ -3305,12 +3496,14 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
targets := state.OutNodes
|
||||
if len(state.ChainHops) > 0 {
|
||||
targets = state.ChainHops[0]
|
||||
} else {
|
||||
targets = h.orderBestExitTargets(state.TunnelID, inNode.NodeID, targets)
|
||||
}
|
||||
chainData, err := buildTunnelChainConfig(state.TunnelID, inNode.NodeID, targets, state.Nodes, state.IPPreference)
|
||||
if err != nil {
|
||||
return createdChains, createdServices, err
|
||||
}
|
||||
if _, err := h.sendNodeCommand(inNode.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
if err := h.applyTunnelChainOnNode(inNode.NodeID, chainData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3320,11 +3513,13 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
}
|
||||
|
||||
for i, hop := range state.ChainHops {
|
||||
nextTargets := state.OutNodes
|
||||
if i+1 < len(state.ChainHops) {
|
||||
nextTargets = state.ChainHops[i+1]
|
||||
}
|
||||
for _, chainNode := range hop {
|
||||
nextTargets := state.OutNodes
|
||||
if i+1 < len(state.ChainHops) {
|
||||
nextTargets = state.ChainHops[i+1]
|
||||
} else {
|
||||
nextTargets = h.orderBestExitTargets(state.TunnelID, chainNode.NodeID, nextTargets)
|
||||
}
|
||||
node := state.Nodes[chainNode.NodeID]
|
||||
if node != nil && (node.IsRemote == 1 || node.Status != 1) {
|
||||
continue
|
||||
@@ -3333,7 +3528,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
if err != nil {
|
||||
return createdChains, createdServices, err
|
||||
}
|
||||
if _, err := h.sendNodeCommand(chainNode.NodeID, "AddChains", chainData, true, false); err != nil {
|
||||
if err := h.applyTunnelChainOnNode(chainNode.NodeID, chainData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3342,7 +3537,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
createdChains = append(createdChains, chainNode.NodeID)
|
||||
|
||||
serviceData := buildTunnelChainServiceConfig(state.TunnelID, chainNode, state.Nodes[chainNode.NodeID], len(nextTargets))
|
||||
if err := h.addTunnelServiceOnNode(chainNode.NodeID, state.TunnelID, serviceData); err != nil {
|
||||
if err := h.addTunnelServiceOnNodeWithMode(chainNode.NodeID, state.TunnelID, serviceData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3358,7 +3553,7 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
continue
|
||||
}
|
||||
serviceData := buildTunnelChainServiceConfig(state.TunnelID, outNode, state.Nodes[outNode.NodeID], 1)
|
||||
if err := h.addTunnelServiceOnNode(outNode.NodeID, state.TunnelID, serviceData); err != nil {
|
||||
if err := h.addTunnelServiceOnNodeWithMode(outNode.NodeID, state.TunnelID, serviceData, upsert); err != nil {
|
||||
if shouldDeferTunnelRuntimeApplyError(err) {
|
||||
continue
|
||||
}
|
||||
@@ -3370,6 +3565,151 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
|
||||
return createdChains, createdServices, nil
|
||||
}
|
||||
|
||||
func (h *Handler) applyTunnelChainOnNode(nodeID int64, chainData map[string]interface{}, upsert bool) error {
|
||||
if upsert {
|
||||
return h.upsertTunnelChainOnNode(nodeID, chainData)
|
||||
}
|
||||
_, err := h.sendNodeCommand(nodeID, "AddChains", chainData, true, false)
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *Handler) applyBestExitChainOrder(tunnelID, ownerNodeID int64, outNodes []chainNodeRecord, scores []bestExitCandidateScore, ipPreference string) error {
|
||||
if h == nil {
|
||||
log.Printf("best_exit: invalid chain update context tunnel=%d owner=%d", tunnelID, ownerNodeID)
|
||||
return errors.New("invalid best exit chain update context")
|
||||
}
|
||||
if tunnelID <= 0 || ownerNodeID <= 0 || len(outNodes) == 0 {
|
||||
log.Printf("best_exit: invalid chain update input tunnel=%d owner=%d exits=%d", tunnelID, ownerNodeID, len(outNodes))
|
||||
return fmt.Errorf("invalid best exit chain update input tunnel=%d owner=%d exits=%d", tunnelID, ownerNodeID, len(outNodes))
|
||||
}
|
||||
targets := chainRecordsToRuntimeTargets(outNodes)
|
||||
orderedIDs := make([]int64, 0, len(scores))
|
||||
for _, score := range scores {
|
||||
if score.ExitNodeID > 0 {
|
||||
orderedIDs = append(orderedIDs, score.ExitNodeID)
|
||||
}
|
||||
}
|
||||
targets = orderRuntimeTargetsByNodeID(targets, orderedIDs)
|
||||
nodes := make(map[int64]*nodeRecord, len(targets)+1)
|
||||
if owner, err := h.getNodeRecord(ownerNodeID); err == nil && owner != nil {
|
||||
nodes[ownerNodeID] = owner
|
||||
}
|
||||
for _, target := range targets {
|
||||
if node, err := h.getNodeRecord(target.NodeID); err == nil && node != nil {
|
||||
nodes[target.NodeID] = node
|
||||
}
|
||||
}
|
||||
owner := nodes[ownerNodeID]
|
||||
if owner != nil && owner.IsRemote == 1 {
|
||||
if err := h.applyRemoteBestExitChainOrder(tunnelID, ownerNodeID, owner, targets, nodes, ipPreference); err != nil {
|
||||
log.Printf("best_exit: update remote federation chain failed tunnel=%d owner=%d err=%v", tunnelID, ownerNodeID, err)
|
||||
return err
|
||||
}
|
||||
log.Printf("best_exit: updated remote federation chain tunnel=%d owner=%d best_exit=%d", tunnelID, ownerNodeID, targets[0].NodeID)
|
||||
return nil
|
||||
}
|
||||
chainData, err := buildTunnelChainConfig(tunnelID, ownerNodeID, targets, nodes, ipPreference)
|
||||
if err != nil {
|
||||
log.Printf("best_exit: build chain failed tunnel=%d owner=%d err=%v", tunnelID, ownerNodeID, err)
|
||||
return err
|
||||
}
|
||||
if err := h.applyTunnelChainOnNode(ownerNodeID, chainData, true); err != nil {
|
||||
log.Printf("best_exit: update chain failed tunnel=%d owner=%d err=%v", tunnelID, ownerNodeID, err)
|
||||
return err
|
||||
}
|
||||
log.Printf("best_exit: updated chain tunnel=%d owner=%d best_exit=%d", tunnelID, ownerNodeID, targets[0].NodeID)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) applyRemoteBestExitChainOrder(tunnelID, ownerNodeID int64, owner *nodeRecord, targets []tunnelRuntimeNode, nodes map[int64]*nodeRecord, ipPreference string) error {
|
||||
if h == nil || h.repo == nil || owner == nil {
|
||||
return errors.New("invalid remote best exit update context")
|
||||
}
|
||||
bindings, err := h.repo.ListActiveFederationTunnelBindingsByTunnel(tunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var binding *repo.FederationTunnelBinding
|
||||
for i := range bindings {
|
||||
if bindings[i].NodeID == ownerNodeID && bindings[i].ChainType == 2 && bindings[i].Status == 1 {
|
||||
binding = &bindings[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
if binding == nil {
|
||||
return fmt.Errorf("active federation middle binding not found for tunnel=%d owner=%d", tunnelID, ownerNodeID)
|
||||
}
|
||||
|
||||
remoteURL := strings.TrimSpace(owner.RemoteURL)
|
||||
if remoteURL == "" {
|
||||
remoteURL = strings.TrimSpace(binding.RemoteURL)
|
||||
}
|
||||
remoteToken := strings.TrimSpace(owner.RemoteToken)
|
||||
if remoteURL == "" || remoteToken == "" {
|
||||
return errors.New("远程节点缺少共享配置")
|
||||
}
|
||||
|
||||
applyTargets := make([]client.RuntimeTarget, 0, len(targets))
|
||||
for _, target := range targets {
|
||||
targetNode := nodes[target.NodeID]
|
||||
if targetNode == nil {
|
||||
return errors.New("节点不存在")
|
||||
}
|
||||
host, hostErr := selectTunnelDialHost(owner, targetNode, ipPreference, target.ConnectIP)
|
||||
if hostErr != nil {
|
||||
return hostErr
|
||||
}
|
||||
if target.Port <= 0 {
|
||||
return errors.New("节点端口不能为空")
|
||||
}
|
||||
applyTargets = append(applyTargets, client.RuntimeTarget{
|
||||
Host: host,
|
||||
Port: target.Port,
|
||||
Protocol: defaultString(target.Protocol, "tls"),
|
||||
})
|
||||
}
|
||||
|
||||
ownerRuntimeNode := tunnelRuntimeNode{NodeID: ownerNodeID, Protocol: "tls", Strategy: "round", ChainType: 2}
|
||||
if chainRows, listErr := h.repo.ListChainNodesForTunnel(tunnelID); listErr == nil {
|
||||
for _, row := range chainRows {
|
||||
if row.NodeID == ownerNodeID && row.ChainType == 2 {
|
||||
ownerRuntimeNode = tunnelRuntimeNode{
|
||||
NodeID: row.NodeID,
|
||||
Protocol: row.Protocol,
|
||||
Strategy: row.Strategy,
|
||||
Inx: int(row.Inx),
|
||||
ChainType: row.ChainType,
|
||||
Port: row.Port,
|
||||
ConnectIP: row.ConnectIP,
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
_, err = client.NewFederationClient().ApplyRole(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeApplyRoleRequest{
|
||||
ResourceKey: strings.TrimSpace(binding.ResourceKey),
|
||||
Role: "middle",
|
||||
Protocol: defaultString(ownerRuntimeNode.Protocol, "tls"),
|
||||
Strategy: runtimeStrategyForTargets(ownerRuntimeNode, targets),
|
||||
Targets: applyTargets,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *Handler) upsertTunnelChainOnNode(nodeID int64, chainData map[string]interface{}) error {
|
||||
if h == nil {
|
||||
return errors.New("invalid tunnel chain context")
|
||||
}
|
||||
chainName := asString(chainData["name"])
|
||||
if strings.TrimSpace(chainName) == "" {
|
||||
return errors.New("转发链名称不能为空")
|
||||
}
|
||||
payload := map[string]interface{}{"chain": chainName, "data": chainData}
|
||||
_, err := h.sendNodeCommand(nodeID, "UpdateChains", payload, true, false)
|
||||
return err
|
||||
}
|
||||
|
||||
func retryTunnelServiceAddWithCleanup(add func() error, cleanup func() error, wait time.Duration) error {
|
||||
if add == nil {
|
||||
return errors.New("invalid tunnel service add callback")
|
||||
@@ -3391,6 +3731,10 @@ func retryTunnelServiceAddWithCleanup(add func() error, cleanup func() error, wa
|
||||
}
|
||||
|
||||
func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []map[string]interface{}) error {
|
||||
return h.addTunnelServiceOnNodeWithMode(nodeID, tunnelID, serviceData, false)
|
||||
}
|
||||
|
||||
func (h *Handler) addTunnelServiceOnNodeWithMode(nodeID, tunnelID int64, serviceData []map[string]interface{}, upsert bool) error {
|
||||
if h == nil {
|
||||
return errors.New("invalid tunnel service context")
|
||||
}
|
||||
@@ -3400,9 +3744,13 @@ func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []m
|
||||
serviceName = strings.TrimSpace(name)
|
||||
}
|
||||
}
|
||||
command := "AddService"
|
||||
if upsert {
|
||||
command = "UpdateService"
|
||||
}
|
||||
return retryTunnelServiceAddWithCleanup(
|
||||
func() error {
|
||||
_, err := h.sendNodeCommand(nodeID, "AddService", serviceData, true, false)
|
||||
_, err := h.sendNodeCommand(nodeID, command, serviceData, true, false)
|
||||
return err
|
||||
},
|
||||
func() error {
|
||||
@@ -3421,16 +3769,7 @@ func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tu
|
||||
protocol = "tls"
|
||||
}
|
||||
seenServices := make(map[int64]struct{})
|
||||
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),
|
||||
}
|
||||
serviceNames := tunnelRuntimeServiceNames(tunnelID)
|
||||
for i := len(serviceNodeIDs) - 1; i >= 0; i-- {
|
||||
nodeID := serviceNodeIDs[i]
|
||||
if _, ok := seenServices[nodeID]; ok {
|
||||
@@ -3528,7 +3867,7 @@ func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRu
|
||||
})
|
||||
}
|
||||
|
||||
strategy := defaultString(strings.TrimSpace(targets[0].Strategy), "round")
|
||||
strategy := runtimeTunnelStrategy(defaultString(strings.TrimSpace(targets[0].Strategy), "round"))
|
||||
hop := map[string]interface{}{
|
||||
"name": fmt.Sprintf("hop_%d", tunnelID),
|
||||
"selector": map[string]interface{}{
|
||||
@@ -3548,6 +3887,24 @@ func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRu
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) orderBestExitTargets(tunnelID, ownerNodeID int64, targets []tunnelRuntimeNode) []tunnelRuntimeNode {
|
||||
if len(targets) <= 1 || !isBestTunnelStrategy(targets[0].Strategy) {
|
||||
return append([]tunnelRuntimeNode(nil), targets...)
|
||||
}
|
||||
if h == nil || h.bestExit == nil {
|
||||
return append([]tunnelRuntimeNode(nil), targets...)
|
||||
}
|
||||
return h.bestExit.orderTargets(bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: ownerNodeID}, targets)
|
||||
}
|
||||
|
||||
func runtimeStrategyForTargets(owner tunnelRuntimeNode, targets []tunnelRuntimeNode) string {
|
||||
strategy := defaultString(owner.Strategy, "round")
|
||||
if len(targets) > 0 {
|
||||
strategy = defaultString(targets[0].Strategy, strategy)
|
||||
}
|
||||
return runtimeTunnelStrategy(strategy)
|
||||
}
|
||||
|
||||
func buildTunnelChainServiceConfig(tunnelID int64, chainNode tunnelRuntimeNode, node *nodeRecord, nextHopCandidateCount int) []map[string]interface{} {
|
||||
if node == nil {
|
||||
return nil
|
||||
@@ -4085,7 +4442,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
||||
h.repo.RollbackForwardFields(
|
||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||
oldForward.SpeedID, oldForward.MaxConn,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
@@ -4249,6 +4606,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() {
|
||||
return fmt.Errorf("address resolves to internal IP: %s", ip.String())
|
||||
return fmt.Errorf("address %q resolves to internal IP: %s", addr, ip.String())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func (h *Handler) storageSummary(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if h == nil || h.repo == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "repository not initialized"))
|
||||
return
|
||||
}
|
||||
|
||||
summary, err := h.repo.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OK(summary))
|
||||
}
|
||||
@@ -0,0 +1,592 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
const (
|
||||
panelDeployDirEnv = "PANEL_DEPLOY_DIR"
|
||||
panelBackendContainerEnv = "PANEL_BACKEND_CONTAINER"
|
||||
defaultPanelDeployDir = "/opt/flvx-panel"
|
||||
defaultPanelBackendName = "flux-panel-backend"
|
||||
dockerSocketPath = "/var/run/docker.sock"
|
||||
maxSystemUpgradeComposeAssetBytes = 1 << 20
|
||||
systemUpgradeMessage = "升级 helper 已启动,面板服务将短暂重启"
|
||||
systemUpgradeConflictError = "已有面板升级任务执行中"
|
||||
)
|
||||
|
||||
var safeBackendContainerPattern = regexp.MustCompile(`^[A-Za-z0-9_.-]+$`)
|
||||
var enableIPv6ComposePattern = regexp.MustCompile(`(?im)^\s*enable_ipv6\s*:\s*['"]?true['"]?\s*(?:#.*)?$`)
|
||||
var systemUpgradeReleaseBaseURL = githubHTMLBase
|
||||
|
||||
type systemUpgradeExecutor struct {
|
||||
deployDir string
|
||||
backendContainer string
|
||||
}
|
||||
|
||||
type systemUpgradeCapabilityData struct {
|
||||
Capable bool `json:"capable"`
|
||||
Reasons []string `json:"reasons"`
|
||||
DeployDir string `json:"deployDir"`
|
||||
BackendContainer string `json:"backendContainer"`
|
||||
}
|
||||
|
||||
type systemUpgradeReleaseData struct {
|
||||
Version string `json:"version"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"publishedAt"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
|
||||
type systemUpgradeVersionData struct {
|
||||
CurrentVersion string `json:"currentVersion"`
|
||||
LatestVersion string `json:"latestVersion"`
|
||||
HasUpdate bool `json:"hasUpdate"`
|
||||
Channel string `json:"channel"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
Capability systemUpgradeCapabilityData `json:"capability"`
|
||||
}
|
||||
|
||||
type systemUpgradeCheckData struct {
|
||||
CurrentVersion string `json:"currentVersion"`
|
||||
LatestVersion string `json:"latestVersion"`
|
||||
HasUpdate bool `json:"hasUpdate"`
|
||||
Channel string `json:"channel"`
|
||||
Capability systemUpgradeCapabilityData `json:"capability"`
|
||||
Releases []systemUpgradeReleaseData `json:"releases"`
|
||||
}
|
||||
|
||||
type systemUpgradeRunData struct {
|
||||
Version string `json:"version"`
|
||||
Channel string `json:"channel"`
|
||||
ComposeAsset string `json:"composeAsset"`
|
||||
HelperContainer string `json:"helperContainer"`
|
||||
BackendImageID string `json:"backendImageId"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
type systemUpgradeRequest struct {
|
||||
Version string `json:"version"`
|
||||
Channel string `json:"channel"`
|
||||
}
|
||||
|
||||
func newSystemUpgradeExecutor() *systemUpgradeExecutor {
|
||||
deployDir := strings.TrimSpace(os.Getenv(panelDeployDirEnv))
|
||||
if deployDir == "" {
|
||||
deployDir = defaultPanelDeployDir
|
||||
}
|
||||
backendContainer := strings.TrimSpace(os.Getenv(panelBackendContainerEnv))
|
||||
if backendContainer == "" {
|
||||
backendContainer = defaultPanelBackendName
|
||||
}
|
||||
return &systemUpgradeExecutor{deployDir: deployDir, backendContainer: backendContainer}
|
||||
}
|
||||
|
||||
func currentPanelVersion() string {
|
||||
version := strings.TrimSpace(os.Getenv("FLUX_VERSION"))
|
||||
if version == "" {
|
||||
return "dev"
|
||||
}
|
||||
return version
|
||||
}
|
||||
|
||||
func validateBackendContainerName(value string) error {
|
||||
if value == "" {
|
||||
return fmt.Errorf("backend container name is empty")
|
||||
}
|
||||
if !safeBackendContainerPattern.MatchString(value) {
|
||||
return fmt.Errorf("unsafe backend container name: %s", value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUpgradeVersion(value string) error {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return fmt.Errorf("upgrade version is empty")
|
||||
}
|
||||
for _, r := range value {
|
||||
if r < 0x20 || r == 0x7f {
|
||||
return fmt.Errorf("unsafe upgrade version: contains control character")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) composePath() string {
|
||||
return filepath.Join(e.deployDir, "docker-compose.yml")
|
||||
}
|
||||
func (e *systemUpgradeExecutor) envPath() string { return filepath.Join(e.deployDir, ".env") }
|
||||
|
||||
func (e *systemUpgradeExecutor) capability(ctx context.Context) systemUpgradeCapabilityData {
|
||||
reasons := make([]string, 0)
|
||||
if !filepath.IsAbs(e.deployDir) {
|
||||
reasons = append(reasons, "部署目录必须是绝对路径")
|
||||
}
|
||||
if err := validateBackendContainerName(e.backendContainer); err != nil {
|
||||
reasons = append(reasons, err.Error())
|
||||
}
|
||||
if out, err := exec.CommandContext(ctx, "docker", "--version").CombinedOutput(); err != nil {
|
||||
reasons = append(reasons, fmt.Sprintf("docker CLI不可用: %v: %s", err, strings.TrimSpace(string(out))))
|
||||
}
|
||||
if info, err := os.Stat(dockerSocketPath); err != nil {
|
||||
reasons = append(reasons, "docker socket不可用: "+err.Error())
|
||||
} else if info.IsDir() {
|
||||
reasons = append(reasons, "docker socket路径不是文件")
|
||||
}
|
||||
if info, err := os.Stat(e.composePath()); err != nil {
|
||||
reasons = append(reasons, "部署docker-compose.yml不可用: "+err.Error())
|
||||
} else if info.IsDir() {
|
||||
reasons = append(reasons, "部署docker-compose.yml不是文件")
|
||||
}
|
||||
if info, err := os.Stat(e.envPath()); err != nil {
|
||||
reasons = append(reasons, "部署.env不可用: "+err.Error())
|
||||
} else if info.IsDir() {
|
||||
reasons = append(reasons, "部署.env不是文件")
|
||||
}
|
||||
if out, err := exec.CommandContext(ctx, "docker", "compose", "version").CombinedOutput(); err != nil {
|
||||
reasons = append(reasons, fmt.Sprintf("docker compose不可用: %v: %s", err, strings.TrimSpace(string(out))))
|
||||
}
|
||||
if _, err := e.currentBackendImage(ctx); err != nil {
|
||||
reasons = append(reasons, err.Error())
|
||||
}
|
||||
|
||||
return systemUpgradeCapabilityData{
|
||||
Capable: len(reasons) == 0,
|
||||
Reasons: reasons,
|
||||
DeployDir: e.deployDir,
|
||||
BackendContainer: e.backendContainer,
|
||||
}
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) selectComposeAsset(current []byte) string {
|
||||
if enableIPv6ComposePattern.Match(current) {
|
||||
return "docker-compose-v6.yml"
|
||||
}
|
||||
return "docker-compose-v4.yml"
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) helperScript() string {
|
||||
return `set -eu
|
||||
LOGFILE="$PANEL_DEPLOY_DIR/upgrade.log"
|
||||
log() { echo "[$(date '+%Y-%m-%d %H:%M:%S')] $*" | tee -a "$LOGFILE"; }
|
||||
|
||||
cd "$PANEL_DEPLOY_DIR"
|
||||
echo "" > "$LOGFILE"
|
||||
log "开始面板升级"
|
||||
log "工作目录: $(pwd)"
|
||||
|
||||
if [ ! -f docker-compose.yml ]; then
|
||||
log "错误: docker-compose.yml 不存在"
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -f .env ]; then
|
||||
log "错误: .env 不存在"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log "拉取新镜像..."
|
||||
if ! docker compose pull backend frontend 2>&1 | tee -a "$LOGFILE"; then
|
||||
log "错误: 拉取镜像失败"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log "等待旧容器释放资源..."
|
||||
sleep 3
|
||||
|
||||
log "重启服务(force-recreate)..."
|
||||
if ! docker compose up -d --force-recreate --remove-orphans backend frontend 2>&1 | tee -a "$LOGFILE"; then
|
||||
log "错误: 重启服务失败"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log "升级完成"
|
||||
`
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) buildHelperRunArgs(imageID, helperName string) ([]string, error) {
|
||||
if err := validateBackendContainerName(e.backendContainer); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []string{
|
||||
"run", "-d", "--rm", "--name", helperName,
|
||||
"--volumes-from", e.backendContainer,
|
||||
"-v", dockerSocketPath + ":" + dockerSocketPath,
|
||||
"-e", panelDeployDirEnv + "=" + e.deployDir,
|
||||
"--entrypoint", "/bin/sh", imageID,
|
||||
"-c", e.helperScript(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) updateEnvVersion(envPath, version string) error {
|
||||
if err := validateUpgradeVersion(version); err != nil {
|
||||
return err
|
||||
}
|
||||
mode, err := fileModeOrDefault(envPath, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
lines := strings.Split(string(data), "\n")
|
||||
replaced := false
|
||||
for i, line := range lines {
|
||||
if strings.HasPrefix(line, "FLUX_VERSION=") {
|
||||
lines[i] = "FLUX_VERSION=" + version
|
||||
replaced = true
|
||||
}
|
||||
}
|
||||
if !replaced {
|
||||
trimmed := strings.TrimRight(strings.Join(lines, "\n"), "\n")
|
||||
if trimmed == "" {
|
||||
trimmed = "FLUX_VERSION=" + version
|
||||
} else {
|
||||
trimmed += "\nFLUX_VERSION=" + version
|
||||
}
|
||||
return writeFileWithMode(envPath, []byte(trimmed+"\n"), mode)
|
||||
}
|
||||
content := strings.TrimRight(strings.Join(lines, "\n"), "\n") + "\n"
|
||||
return writeFileWithMode(envPath, []byte(content), mode)
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) backupFile(path string) (string, error) {
|
||||
mode, err := fileModeOrDefault(path, 0o600)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
backupPath := path + ".upgrade.bak"
|
||||
if err := writeFileWithMode(backupPath, data, mode); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return backupPath, nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) restoreBackup(path string) error {
|
||||
backupPath := path + ".upgrade.bak"
|
||||
mode, err := fileModeOrDefault(backupPath, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := os.ReadFile(backupPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeFileWithMode(path, data, mode)
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) restoreUpgradeBackups(paths ...string) error {
|
||||
var errs []string
|
||||
for _, path := range paths {
|
||||
if err := e.restoreBackup(path); err != nil {
|
||||
errs = append(errs, fmt.Sprintf("%s: %v", path, err))
|
||||
}
|
||||
}
|
||||
if len(errs) > 0 {
|
||||
return fmt.Errorf("%s", strings.Join(errs, "; "))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) replaceCompose(path string, data []byte) error {
|
||||
if len(bytes.TrimSpace(data)) == 0 {
|
||||
return fmt.Errorf("compose asset is empty")
|
||||
}
|
||||
mode, err := fileModeOrDefault(path, 0o644)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeFileWithMode(path, data, mode)
|
||||
}
|
||||
|
||||
func fileModeOrDefault(path string, fallback os.FileMode) (os.FileMode, error) {
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return fallback, nil
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
return info.Mode().Perm(), nil
|
||||
}
|
||||
|
||||
func writeFileWithMode(path string, data []byte, mode os.FileMode) error {
|
||||
if err := os.WriteFile(path, data, mode); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Chmod(path, mode)
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) currentBackendImage(ctx context.Context) (string, error) {
|
||||
if err := validateBackendContainerName(e.backendContainer); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out, err := exec.CommandContext(ctx, "docker", "inspect", "-f", "{{.Image}}", e.backendContainer).CombinedOutput()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("inspect backend image failed: %v: %s", err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
imageID := strings.TrimSpace(string(out))
|
||||
if imageID == "" {
|
||||
return "", fmt.Errorf("backend image id is empty")
|
||||
}
|
||||
return imageID, nil
|
||||
}
|
||||
|
||||
func (e *systemUpgradeExecutor) startHelper(ctx context.Context, imageID, helperName string) (string, error) {
|
||||
args, err := e.buildHelperRunArgs(imageID, helperName)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
out, err := exec.CommandContext(ctx, "docker", args...).CombinedOutput()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("start helper failed: %v: %s", err, strings.TrimSpace(string(out)))
|
||||
}
|
||||
containerID := strings.TrimSpace(string(out))
|
||||
if containerID == "" {
|
||||
containerID = helperName
|
||||
}
|
||||
return containerID, nil
|
||||
}
|
||||
|
||||
func (h *Handler) downloadReleaseAsset(version, filename string) ([]byte, error) {
|
||||
url := fmt.Sprintf("%s/%s/releases/download/%s/%s", strings.TrimRight(systemUpgradeReleaseBaseURL, "/"), githubRepo, version, filename)
|
||||
client := &http.Client{Timeout: 60 * time.Second}
|
||||
resp, err := client.Get(url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("下载%s失败: %v", filename, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||
return nil, fmt.Errorf("下载%s返回 %d: %s", filename, resp.StatusCode, strings.TrimSpace(string(body)))
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxSystemUpgradeComposeAssetBytes+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取%s失败: %v", filename, err)
|
||||
}
|
||||
if len(body) > maxSystemUpgradeComposeAssetBytes {
|
||||
return nil, fmt.Errorf("下载%s过大", filename)
|
||||
}
|
||||
if len(bytes.TrimSpace(body)) == 0 {
|
||||
return nil, fmt.Errorf("下载%s内容为空", filename)
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func releasesForChannel(releases []githubRelease, channel string) []systemUpgradeReleaseData {
|
||||
channel = normalizeReleaseChannel(channel)
|
||||
items := make([]systemUpgradeReleaseData, 0, len(releases))
|
||||
for _, r := range releases {
|
||||
if r.Draft {
|
||||
continue
|
||||
}
|
||||
tag := strings.TrimSpace(r.TagName)
|
||||
if tag == "" {
|
||||
continue
|
||||
}
|
||||
itemChannel := releaseChannelFromTag(tag)
|
||||
if itemChannel != channel {
|
||||
continue
|
||||
}
|
||||
items = append(items, systemUpgradeReleaseData{
|
||||
Version: tag,
|
||||
Name: r.Name,
|
||||
PublishedAt: r.PublishedAt,
|
||||
Prerelease: itemChannel == releaseChannelDev,
|
||||
Channel: itemChannel,
|
||||
})
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func decodeSystemUpgradeRequest(r *http.Request, req *systemUpgradeRequest) error {
|
||||
defer r.Body.Close()
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(bytes.TrimSpace(body)) == 0 {
|
||||
return nil
|
||||
}
|
||||
decoder := json.NewDecoder(bytes.NewReader(body))
|
||||
decoder.DisallowUnknownFields()
|
||||
return decoder.Decode(req)
|
||||
}
|
||||
|
||||
func systemUpgradeVersionResponse(current, channel, latest string, lookupErr error, capability systemUpgradeCapabilityData) systemUpgradeVersionData {
|
||||
data := systemUpgradeVersionData{
|
||||
CurrentVersion: current,
|
||||
LatestVersion: latest,
|
||||
HasUpdate: latest != "" && latest != current,
|
||||
Channel: channel,
|
||||
Capability: capability,
|
||||
}
|
||||
if lookupErr != nil {
|
||||
data.LatestVersion = ""
|
||||
data.HasUpdate = false
|
||||
data.Reason = lookupErr.Error()
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func (h *Handler) systemVersion(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
channel := releaseChannelStable
|
||||
current := currentPanelVersion()
|
||||
exec := newSystemUpgradeExecutor()
|
||||
capability := exec.capability(r.Context())
|
||||
latest, err := resolveLatestReleaseByChannel(channel)
|
||||
response.WriteJSON(w, response.OK(systemUpgradeVersionResponse(current, channel, latest, err, capability)))
|
||||
}
|
||||
|
||||
func (h *Handler) systemCheckUpdates(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req systemUpgradeRequest
|
||||
if err := decodeSystemUpgradeRequest(r, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
current := currentPanelVersion()
|
||||
exec := newSystemUpgradeExecutor()
|
||||
capability := exec.capability(r.Context())
|
||||
|
||||
githubReleases, err := fetchGitHubReleases(50)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
releases := releasesForChannel(githubReleases, channel)
|
||||
latest := ""
|
||||
if len(releases) > 0 {
|
||||
latest = releases[0].Version
|
||||
}
|
||||
response.WriteJSON(w, response.OK(systemUpgradeCheckData{
|
||||
CurrentVersion: current,
|
||||
LatestVersion: latest,
|
||||
HasUpdate: latest != "" && latest != current,
|
||||
Channel: channel,
|
||||
Capability: capability,
|
||||
Releases: releases,
|
||||
}))
|
||||
}
|
||||
|
||||
func (h *Handler) systemUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
if !h.systemUpgradeMu.TryLock() {
|
||||
response.WriteJSON(w, response.ErrDefault(systemUpgradeConflictError))
|
||||
return
|
||||
}
|
||||
defer h.systemUpgradeMu.Unlock()
|
||||
|
||||
var req systemUpgradeRequest
|
||||
if err := decodeSystemUpgradeRequest(r, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
channel := normalizeReleaseChannel(req.Channel)
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestReleaseByChannel(channel)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新%s失败: %v", releaseChannelLabel(channel), err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
exec := newSystemUpgradeExecutor()
|
||||
capability := exec.capability(r.Context())
|
||||
if !capability.Capable {
|
||||
response.WriteJSON(w, response.ErrDefault("当前环境不支持面板自升级: "+strings.Join(capability.Reasons, "; ")))
|
||||
return
|
||||
}
|
||||
imageID, err := exec.currentBackendImage(r.Context())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
composePath := exec.composePath()
|
||||
envPath := exec.envPath()
|
||||
composeData, err := os.ReadFile(composePath)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "读取compose失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
composeAsset := exec.selectComposeAsset(composeData)
|
||||
newCompose, err := h.downloadReleaseAsset(version, composeAsset)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err := exec.backupFile(composePath); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "备份compose失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
if _, err := exec.backupFile(envPath); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, "备份.env失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
if err := exec.replaceCompose(composePath, newCompose); err != nil {
|
||||
if restoreErr := exec.restoreUpgradeBackups(composePath, envPath); restoreErr != nil {
|
||||
err = fmt.Errorf("%v; 回滚失败: %v", err, restoreErr)
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, "替换compose失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
if err := exec.updateEnvVersion(envPath, version); err != nil {
|
||||
if restoreErr := exec.restoreUpgradeBackups(composePath, envPath); restoreErr != nil {
|
||||
err = fmt.Errorf("%v; 回滚失败: %v", err, restoreErr)
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, "更新版本配置失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
helperName := fmt.Sprintf("flvx-upgrade-helper-%d", time.Now().Unix())
|
||||
helperContainer, err := exec.startHelper(r.Context(), imageID, helperName)
|
||||
if err != nil {
|
||||
if restoreErr := exec.restoreUpgradeBackups(composePath, envPath); restoreErr != nil {
|
||||
err = fmt.Errorf("%v; 回滚失败: %v", err, restoreErr)
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(systemUpgradeRunData{
|
||||
Version: version,
|
||||
Channel: channel,
|
||||
ComposeAsset: composeAsset,
|
||||
HelperContainer: helperContainer,
|
||||
BackendImageID: imageID,
|
||||
Message: systemUpgradeMessage,
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,398 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSelectComposeAssetUsesIPv6Template(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
compose := []byte("networks:\n gost-network:\n enable_ipv6: true\n")
|
||||
|
||||
if got := exec.selectComposeAsset(compose); got != "docker-compose-v6.yml" {
|
||||
t.Fatalf("selectComposeAsset() = %q, want %q", got, "docker-compose-v6.yml")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadReleaseAssetUsesDirectReleaseURL(t *testing.T) {
|
||||
var gotPath string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
gotPath = r.URL.Path
|
||||
_, _ = w.Write([]byte("services:\n backend:\n image: test\n"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
originalBase := systemUpgradeReleaseBaseURL
|
||||
systemUpgradeReleaseBaseURL = server.URL
|
||||
t.Cleanup(func() { systemUpgradeReleaseBaseURL = originalBase })
|
||||
|
||||
h := &Handler{}
|
||||
data, err := h.downloadReleaseAsset("2.1.9", "docker-compose-v4.yml")
|
||||
if err != nil {
|
||||
t.Fatalf("downloadReleaseAsset() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(string(data), "backend") {
|
||||
t.Fatalf("downloadReleaseAsset() data = %q, want compose data", string(data))
|
||||
}
|
||||
|
||||
wantPath := "/" + githubRepo + "/releases/download/2.1.9/docker-compose-v4.yml"
|
||||
if gotPath != wantPath {
|
||||
t.Fatalf("download path = %q, want %q", gotPath, wantPath)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadReleaseAssetRejectsOversizedBody(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = w.Write(bytes.Repeat([]byte("a"), maxSystemUpgradeComposeAssetBytes+1))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
originalBase := systemUpgradeReleaseBaseURL
|
||||
systemUpgradeReleaseBaseURL = server.URL
|
||||
t.Cleanup(func() { systemUpgradeReleaseBaseURL = originalBase })
|
||||
|
||||
h := &Handler{}
|
||||
_, err := h.downloadReleaseAsset("2.1.9", "docker-compose-v4.yml")
|
||||
if err == nil || !strings.Contains(err.Error(), "过大") {
|
||||
t.Fatalf("downloadReleaseAsset() error = %v, want oversized error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectComposeAssetUsesIPv6TemplateForYAMLVariants(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
for _, compose := range [][]byte{
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6:true\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: True\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: \"true\"\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: 'true'\n"),
|
||||
[]byte("networks:\n gost-network:\n enable_ipv6: true # comment\n"),
|
||||
} {
|
||||
if got := exec.selectComposeAsset(compose); got != "docker-compose-v6.yml" {
|
||||
t.Fatalf("selectComposeAsset(%q) = %q, want %q", string(compose), got, "docker-compose-v6.yml")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSelectComposeAssetFallsBackToIPv4Template(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
compose := []byte("services:\n backend:\n image: test\n")
|
||||
|
||||
if got := exec.selectComposeAsset(compose); got != "docker-compose-v4.yml" {
|
||||
t.Fatalf("selectComposeAsset() = %q, want %q", got, "docker-compose-v4.yml")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionReplacesExistingValue(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, "2.1.9"); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
|
||||
want := "FLUX_VERSION=2.1.9\nJWT_SECRET=test\n"
|
||||
if string(data) != want {
|
||||
t.Fatalf("env content = %q, want %q", string(data), want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionAppendsMissingValue(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("JWT_SECRET=test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, "2.1.9"); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
|
||||
want := "JWT_SECRET=test\nFLUX_VERSION=2.1.9\n"
|
||||
if string(data) != want {
|
||||
t.Fatalf("env content = %q, want %q", string(data), want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionRejectsUnsafeValue(t *testing.T) {
|
||||
for _, version := range []string{"", "2.1.9\nJWT_SECRET=bad", "2.1.9\rbad", "2.1.9\x00bad", "2.1.9\x1fbad"} {
|
||||
t.Run(version, func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
original := []byte("JWT_SECRET=test\n")
|
||||
if err := os.WriteFile(envPath, original, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, version); err == nil {
|
||||
t.Fatal("expected unsafe version to fail validation")
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if string(data) != string(original) {
|
||||
t.Fatalf("env content changed to %q, want %q", string(data), string(original))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionAcceptsVersionLabels(t *testing.T) {
|
||||
for _, version := range []string{"2.1.9", "2.1.9-beta14", "v-test"} {
|
||||
t.Run(version, func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("JWT_SECRET=test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, version); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateEnvVersionPreservesFileMode(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(envPath, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if err := exec.updateEnvVersion(envPath, "2.1.9"); err != nil {
|
||||
t.Fatalf("updateEnvVersion() error = %v", err)
|
||||
}
|
||||
|
||||
info, err := os.Stat(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat() error = %v", err)
|
||||
}
|
||||
if got := info.Mode().Perm(); got != 0o600 {
|
||||
t.Fatalf("env mode = %o, want 0600", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateBackendContainerNameRejectsUnsafeValue(t *testing.T) {
|
||||
if err := validateBackendContainerName("flux-panel-backend;rm -rf /"); err == nil {
|
||||
t.Fatal("expected unsafe container name to fail validation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildHelperRunArgsUsesDetachedContainer(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend"}
|
||||
args, err := exec.buildHelperRunArgs("sha256:abc", "flvx-upgrade-helper")
|
||||
if err != nil {
|
||||
t.Fatalf("buildHelperRunArgs() error = %v", err)
|
||||
}
|
||||
want := []string{
|
||||
"run", "-d", "--rm", "--name", "flvx-upgrade-helper",
|
||||
"--volumes-from", "flux-panel-backend",
|
||||
"-v", "/var/run/docker.sock:/var/run/docker.sock",
|
||||
"-e", "PANEL_DEPLOY_DIR=/opt/flvx-panel",
|
||||
"--entrypoint", "/bin/sh", "sha256:abc",
|
||||
"-c", exec.helperScript(),
|
||||
}
|
||||
|
||||
if !reflect.DeepEqual(args, want) {
|
||||
t.Fatalf("buildHelperRunArgs() = %#v, want %#v", args, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildHelperRunArgsRejectsUnsafeBackendContainer(t *testing.T) {
|
||||
exec := &systemUpgradeExecutor{deployDir: "/opt/flvx-panel", backendContainer: "flux-panel-backend;rm -rf /"}
|
||||
if _, err := exec.buildHelperRunArgs("sha256:abc", "flvx-upgrade-helper"); err == nil {
|
||||
t.Fatal("expected unsafe backend container name to fail validation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemVersionRejectsWrongMethod(t *testing.T) {
|
||||
h := &Handler{}
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/system/version", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
h.systemVersion(rr, req)
|
||||
|
||||
if !strings.Contains(rr.Body.String(), "请求失败") {
|
||||
t.Fatalf("expected wrong-method response, got %s", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUpgradeRejectsConcurrentRequests(t *testing.T) {
|
||||
h := &Handler{}
|
||||
h.systemUpgradeMu.Lock()
|
||||
defer h.systemUpgradeMu.Unlock()
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/upgrade", strings.NewReader(`{"channel":"stable"}`))
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
h.systemUpgrade(rr, req)
|
||||
|
||||
if !strings.Contains(rr.Body.String(), systemUpgradeConflictError) {
|
||||
t.Fatalf("expected conflict message, got %s", rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUpgradeFailsFastBeforeMutatingFiles(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
composePath := filepath.Join(dir, "docker-compose.yml")
|
||||
envPath := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(composePath, []byte("services:\n backend:\n image: test\n"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() compose error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(envPath, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() env error = %v", err)
|
||||
}
|
||||
|
||||
fakeDockerDir := t.TempDir()
|
||||
fakeDockerPath := filepath.Join(fakeDockerDir, "docker")
|
||||
fakeDockerScript := "#!/bin/sh\ncase \"$1\" in\n --version)\n echo 'Docker version 27.0.0'\n exit 0\n ;;&\n compose)\n if [ \"$2\" = version ]; then\n echo 'Docker Compose version v2.33.0'\n exit 0\n fi\n exit 0\n ;;&\n inspect)\n echo 'No such object: flux-panel-backend' >&2\n exit 1\n ;;&\n *)\n exit 0\n ;;&\n esac\n"
|
||||
if err := os.WriteFile(fakeDockerPath, []byte(fakeDockerScript), 0o755); err != nil {
|
||||
t.Fatalf("WriteFile() fake docker error = %v", err)
|
||||
}
|
||||
t.Setenv("PATH", fakeDockerDir+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
t.Setenv(panelDeployDirEnv, dir)
|
||||
t.Setenv(panelBackendContainerEnv, "flux-panel-backend")
|
||||
|
||||
h := &Handler{}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/upgrade", strings.NewReader(`{"channel":"stable"}`))
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
h.systemUpgrade(rr, req)
|
||||
|
||||
if !strings.Contains(rr.Body.String(), "当前环境不支持面板自升级") {
|
||||
t.Fatalf("expected fail-fast capability error, got %s", rr.Body.String())
|
||||
}
|
||||
if _, err := os.Stat(composePath + ".upgrade.bak"); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected no compose backup, got err=%v", err)
|
||||
}
|
||||
if _, err := os.Stat(envPath + ".upgrade.bak"); !os.IsNotExist(err) {
|
||||
t.Fatalf("expected no env backup, got err=%v", err)
|
||||
}
|
||||
composeData, err := os.ReadFile(composePath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() compose error = %v", err)
|
||||
}
|
||||
if string(composeData) != "services:\n backend:\n image: test\n" {
|
||||
t.Fatalf("compose mutated unexpectedly: %q", string(composeData))
|
||||
}
|
||||
envData, err := os.ReadFile(envPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() env error = %v", err)
|
||||
}
|
||||
if string(envData) != "FLUX_VERSION=2.1.8\nJWT_SECRET=test\n" {
|
||||
t.Fatalf("env mutated unexpectedly: %q", string(envData))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpgradeBackupUsesStablePathAndRestoreRestoresOriginal(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "docker-compose.yml")
|
||||
if err := os.WriteFile(path, []byte("original"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
backupPath, err := exec.backupFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("backupFile() error = %v", err)
|
||||
}
|
||||
if backupPath != path+".upgrade.bak" {
|
||||
t.Fatalf("backup path = %q, want %q", backupPath, path+".upgrade.bak")
|
||||
}
|
||||
if err := os.WriteFile(path, []byte("mutated"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
if err := exec.restoreBackup(path); err != nil {
|
||||
t.Fatalf("restoreBackup() error = %v", err)
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if string(data) != "original" {
|
||||
t.Fatalf("restored content = %q, want original", string(data))
|
||||
}
|
||||
}
|
||||
|
||||
func TestRestoreBackupPreservesOriginalFileMode(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, ".env")
|
||||
if err := os.WriteFile(path, []byte("FLUX_VERSION=2.1.8\nJWT_SECRET=test\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
exec := &systemUpgradeExecutor{deployDir: dir, backendContainer: "flux-panel-backend"}
|
||||
if _, err := exec.backupFile(path); err != nil {
|
||||
t.Fatalf("backupFile() error = %v", err)
|
||||
}
|
||||
if err := os.Remove(path); err != nil {
|
||||
t.Fatalf("Remove() error = %v", err)
|
||||
}
|
||||
if err := exec.restoreBackup(path); err != nil {
|
||||
t.Fatalf("restoreBackup() error = %v", err)
|
||||
}
|
||||
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat() error = %v", err)
|
||||
}
|
||||
if got := info.Mode().Perm(); got != 0o600 {
|
||||
t.Fatalf("restored mode = %o, want 0600", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSystemUpgradeRequestRejectsTruncatedJSON(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/check-updates", strings.NewReader(`{"channel":"stable"`))
|
||||
var payload systemUpgradeRequest
|
||||
|
||||
if err := decodeSystemUpgradeRequest(req, &payload); err == nil {
|
||||
t.Fatal("expected truncated JSON to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSystemUpgradeRequestAllowsEmptyBody(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/system/check-updates", strings.NewReader(""))
|
||||
var payload systemUpgradeRequest
|
||||
|
||||
if err := decodeSystemUpgradeRequest(req, &payload); err != nil {
|
||||
t.Fatalf("expected empty body to be accepted, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemUpgradeVersionDataSurfacesLookupFailureReason(t *testing.T) {
|
||||
data, err := json.Marshal(systemUpgradeVersionData{Reason: "GitHub unavailable"})
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(string(data), `"reason":"GitHub unavailable"`) {
|
||||
t.Fatalf("expected reason field in JSON, got %s", string(data))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,436 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
tunnelStrategyBest = "best"
|
||||
bestExitRuntimeStrategy = "fifo"
|
||||
bestExitPublicTargetHost = "www.bing.com"
|
||||
bestExitPublicTargetPort = 443
|
||||
bestExitLossPenaltyMsPerPercent = 100.0
|
||||
bestExitConfirmationRounds = 3
|
||||
bestExitSwitchCooldown = 30 * time.Second
|
||||
bestExitApplyRetryCooldown = bestExitSwitchCooldown
|
||||
bestExitMinLatencyAdvantageMs = 20.0
|
||||
bestExitMinScoreAdvantageRatio = 0.15
|
||||
)
|
||||
|
||||
type bestExitOwnerKey struct {
|
||||
TunnelID int64
|
||||
OwnerNodeID int64
|
||||
}
|
||||
|
||||
type bestExitCandidateScore struct {
|
||||
OwnerNodeID int64
|
||||
ExitNodeID int64
|
||||
ExitName string
|
||||
|
||||
OwnerToExitLatency float64
|
||||
ExitToBingLatency float64
|
||||
OwnerToExitLoss float64
|
||||
ExitToBingLoss float64
|
||||
TotalLatency float64
|
||||
TotalLoss float64
|
||||
Score float64
|
||||
Success bool
|
||||
ErrorMessage string
|
||||
}
|
||||
|
||||
type bestExitSwitchDecision struct {
|
||||
Switch bool
|
||||
ExitNodeID int64
|
||||
Reason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitProbeFunc func(nodeID int64, ip string, port int, options diagnosisExecOptions) (latency float64, loss float64, err error)
|
||||
|
||||
type bestExitProbeResult struct {
|
||||
latency float64
|
||||
loss float64
|
||||
err error
|
||||
}
|
||||
|
||||
type bestExitProbeCacheKey struct {
|
||||
NodeID int64
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
type bestExitDecision struct {
|
||||
AppliedExitNodeID int64
|
||||
PendingExitNodeID int64
|
||||
PendingCount int
|
||||
LastSwitchAt time.Time
|
||||
LastApplyFailureAt time.Time
|
||||
LastApplyFailureExitNodeID int64
|
||||
LastReason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitManager struct {
|
||||
mu sync.Mutex
|
||||
decisions map[bestExitOwnerKey]*bestExitDecision
|
||||
}
|
||||
|
||||
func newBestExitManager() *bestExitManager {
|
||||
return &bestExitManager{decisions: make(map[bestExitOwnerKey]*bestExitDecision)}
|
||||
}
|
||||
|
||||
func isBestTunnelStrategy(strategy string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(strategy), tunnelStrategyBest)
|
||||
}
|
||||
|
||||
func runtimeTunnelStrategy(strategy string) string {
|
||||
if isBestTunnelStrategy(strategy) {
|
||||
return bestExitRuntimeStrategy
|
||||
}
|
||||
return strategy
|
||||
}
|
||||
|
||||
func scoreBestExitCandidate(ownerNodeID int64, exit chainNodeRecord, ownerLatency, ownerLoss, publicLatency, publicLoss float64) bestExitCandidateScore {
|
||||
totalLatency := ownerLatency + publicLatency
|
||||
totalLoss := combineLossPercent(ownerLoss, publicLoss)
|
||||
return bestExitCandidateScore{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
ExitNodeID: exit.NodeID,
|
||||
ExitName: exit.NodeName,
|
||||
OwnerToExitLatency: ownerLatency,
|
||||
ExitToBingLatency: publicLatency,
|
||||
OwnerToExitLoss: ownerLoss,
|
||||
ExitToBingLoss: publicLoss,
|
||||
TotalLatency: totalLatency,
|
||||
TotalLoss: totalLoss,
|
||||
Score: totalLatency + totalLoss*bestExitLossPenaltyMsPerPercent,
|
||||
Success: true,
|
||||
}
|
||||
}
|
||||
|
||||
func failedBestExitCandidate(ownerNodeID int64, exit chainNodeRecord, message string) bestExitCandidateScore {
|
||||
return bestExitCandidateScore{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
ExitNodeID: exit.NodeID,
|
||||
ExitName: exit.NodeName,
|
||||
Success: false,
|
||||
ErrorMessage: message,
|
||||
}
|
||||
}
|
||||
|
||||
func combineLossPercent(a, b float64) float64 {
|
||||
a = clampPercent(a)
|
||||
b = clampPercent(b)
|
||||
return (1 - (1-a/100.0)*(1-b/100.0)) * 100.0
|
||||
}
|
||||
|
||||
func clampPercent(v float64) float64 {
|
||||
if v < 0 {
|
||||
return 0
|
||||
}
|
||||
if v > 100 {
|
||||
return 100
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func sortBestExitScores(scores []bestExitCandidateScore) {
|
||||
sort.SliceStable(scores, func(i, j int) bool {
|
||||
return bestExitScoreLess(scores[i], scores[j])
|
||||
})
|
||||
}
|
||||
|
||||
func evaluateBestExitOwner(owner chainNodeRecord, exits []chainNodeRecord, nodes map[int64]*nodeRecord, ipPreference string, options diagnosisExecOptions, target tunnelProbeTarget, ping bestExitProbeFunc) []bestExitCandidateScore {
|
||||
scores := make([]bestExitCandidateScore, 0, len(exits))
|
||||
if owner.NodeID <= 0 || len(exits) == 0 || ping == nil {
|
||||
return scores
|
||||
}
|
||||
ownerNode := nodes[owner.NodeID]
|
||||
for _, exit := range exits {
|
||||
exitNode := nodes[exit.NodeID]
|
||||
if exitNode == nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "exit node unavailable"))
|
||||
continue
|
||||
}
|
||||
targetIP, targetPort, resolveErr := resolveBestExitProbeTarget(ownerNode, exitNode, exit.Port, ipPreference, exit.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, resolveErr.Error()))
|
||||
continue
|
||||
}
|
||||
ownerLatency, ownerLoss, ownerErr := ping(owner.NodeID, targetIP, targetPort, options)
|
||||
if ownerErr != nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, ownerErr.Error()))
|
||||
continue
|
||||
}
|
||||
publicLatency, publicLoss, publicErr := ping(exit.NodeID, target.Host, target.Port, options)
|
||||
if publicErr != nil {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, publicErr.Error()))
|
||||
continue
|
||||
}
|
||||
scores = append(scores, scoreBestExitCandidate(owner.NodeID, exit, ownerLatency, ownerLoss, publicLatency, publicLoss))
|
||||
}
|
||||
sortBestExitScores(scores)
|
||||
return scores
|
||||
}
|
||||
|
||||
func resolveBestExitProbeTarget(fromNode, targetNode *nodeRecord, preferredPort int, ipPreference string, connectIP string) (string, int, error) {
|
||||
if targetNode == nil {
|
||||
return "", 0, errors.New("目标节点不存在")
|
||||
}
|
||||
host, err := selectTunnelDialHost(fromNode, targetNode, ipPreference, connectIP)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
if strings.TrimSpace(host) == "" {
|
||||
return "", 0, errors.New("目标节点地址为空")
|
||||
}
|
||||
port := preferredPort
|
||||
if port <= 0 {
|
||||
port = firstPortFromRange(targetNode.PortRange)
|
||||
}
|
||||
if port <= 0 {
|
||||
port = 443
|
||||
}
|
||||
return host, port, nil
|
||||
}
|
||||
|
||||
func newBestExitRoundPinger(base bestExitProbeFunc) bestExitProbeFunc {
|
||||
cache := make(map[bestExitProbeCacheKey]bestExitProbeResult)
|
||||
return func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
key := bestExitProbeCacheKey{NodeID: nodeID, Host: ip, Port: port}
|
||||
if cached, ok := cache[key]; ok {
|
||||
return cached.latency, cached.loss, cached.err
|
||||
}
|
||||
lat, loss, err := base(nodeID, ip, port, options)
|
||||
cache[key] = bestExitProbeResult{latency: lat, loss: loss, err: err}
|
||||
return lat, loss, err
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitChainOwners(inNodes []chainNodeRecord, chainHops [][]chainNodeRecord) []chainNodeRecord {
|
||||
if len(chainHops) == 0 {
|
||||
return inNodes
|
||||
}
|
||||
return chainHops[len(chainHops)-1]
|
||||
}
|
||||
|
||||
func chainRecordsToRuntimeTargets(rows []chainNodeRecord) []tunnelRuntimeNode {
|
||||
out := make([]tunnelRuntimeNode, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
out = append(out, tunnelRuntimeNode{
|
||||
NodeID: row.NodeID,
|
||||
Protocol: row.Protocol,
|
||||
Strategy: row.Strategy,
|
||||
Inx: int(row.Inx),
|
||||
ChainType: row.ChainType,
|
||||
Port: row.Port,
|
||||
ConnectIP: row.ConnectIP,
|
||||
})
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func orderRuntimeTargetsByNodeID(targets []tunnelRuntimeNode, orderedIDs []int64) []tunnelRuntimeNode {
|
||||
out := append([]tunnelRuntimeNode(nil), targets...)
|
||||
if len(out) <= 1 || len(orderedIDs) == 0 {
|
||||
return out
|
||||
}
|
||||
positions := make(map[int64]int, len(orderedIDs))
|
||||
for i, id := range orderedIDs {
|
||||
if _, ok := positions[id]; !ok {
|
||||
positions[id] = i
|
||||
}
|
||||
}
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
pi, iok := positions[out[i].NodeID]
|
||||
pj, jok := positions[out[j].NodeID]
|
||||
if iok != jok {
|
||||
return iok
|
||||
}
|
||||
if iok && jok && pi != pj {
|
||||
return pi < pj
|
||||
}
|
||||
return false
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func cloneBestExitScores(scores []bestExitCandidateScore) []bestExitCandidateScore {
|
||||
return append([]bestExitCandidateScore(nil), scores...)
|
||||
}
|
||||
|
||||
func bestExitDecisionResult(switchNow bool, exitNodeID int64, reason string, scores []bestExitCandidateScore) bestExitSwitchDecision {
|
||||
return bestExitSwitchDecision{Switch: switchNow, ExitNodeID: exitNodeID, Reason: reason, Scores: cloneBestExitScores(scores)}
|
||||
}
|
||||
|
||||
func bestExitScoreLess(a, b bestExitCandidateScore) bool {
|
||||
if a.Success != b.Success {
|
||||
return a.Success
|
||||
}
|
||||
if !a.Success && !b.Success {
|
||||
return a.ExitNodeID < b.ExitNodeID
|
||||
}
|
||||
if a.Score != b.Score {
|
||||
return a.Score < b.Score
|
||||
}
|
||||
return a.ExitNodeID < b.ExitNodeID
|
||||
}
|
||||
|
||||
func bestExitHasMinimumAdvantage(candidate, current bestExitCandidateScore) bool {
|
||||
if !candidate.Success {
|
||||
return false
|
||||
}
|
||||
if !current.Success {
|
||||
return true
|
||||
}
|
||||
improvement := current.Score - candidate.Score
|
||||
threshold := current.Score * bestExitMinScoreAdvantageRatio
|
||||
if threshold < bestExitMinLatencyAdvantageMs {
|
||||
threshold = bestExitMinLatencyAdvantageMs
|
||||
}
|
||||
return improvement >= threshold
|
||||
}
|
||||
|
||||
func (m *bestExitManager) setApplied(key bestExitOwnerKey, exitNodeID int64, at time.Time) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisionLocked(key)
|
||||
d.AppliedExitNodeID = exitNodeID
|
||||
d.PendingExitNodeID = 0
|
||||
d.PendingCount = 0
|
||||
d.LastApplyFailureAt = time.Time{}
|
||||
d.LastApplyFailureExitNodeID = 0
|
||||
d.LastSwitchAt = at
|
||||
}
|
||||
|
||||
func (m *bestExitManager) recordApplyFailure(key bestExitOwnerKey, exitNodeID int64, at time.Time) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisionLocked(key)
|
||||
d.LastApplyFailureAt = at
|
||||
d.LastApplyFailureExitNodeID = exitNodeID
|
||||
d.LastReason = "apply retry cooldown"
|
||||
}
|
||||
|
||||
func (m *bestExitManager) ensureApplied(key bestExitOwnerKey, exitNodeID int64, at time.Time) {
|
||||
if m == nil || exitNodeID <= 0 {
|
||||
return
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisionLocked(key)
|
||||
if d.AppliedExitNodeID == 0 {
|
||||
d.AppliedExitNodeID = exitNodeID
|
||||
d.LastSwitchAt = at
|
||||
}
|
||||
}
|
||||
|
||||
func (m *bestExitManager) observeScores(key bestExitOwnerKey, scores []bestExitCandidateScore, now time.Time) bestExitSwitchDecision {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
ordered := append([]bestExitCandidateScore(nil), scores...)
|
||||
sortBestExitScores(ordered)
|
||||
d := m.decisionLocked(key)
|
||||
d.Scores = cloneBestExitScores(ordered)
|
||||
|
||||
if len(ordered) == 0 || !ordered[0].Success {
|
||||
d.LastReason = "all exits failed"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
candidate := ordered[0]
|
||||
if d.AppliedExitNodeID == 0 {
|
||||
d.AppliedExitNodeID = candidate.ExitNodeID
|
||||
d.LastSwitchAt = now
|
||||
d.LastReason = "initial best exit"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
if candidate.ExitNodeID == d.AppliedExitNodeID {
|
||||
d.PendingExitNodeID = 0
|
||||
d.PendingCount = 0
|
||||
d.LastApplyFailureAt = time.Time{}
|
||||
d.LastApplyFailureExitNodeID = 0
|
||||
d.LastReason = "current exit remains best"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
if candidate.ExitNodeID == d.LastApplyFailureExitNodeID && !d.LastApplyFailureAt.IsZero() && now.Sub(d.LastApplyFailureAt) < bestExitApplyRetryCooldown {
|
||||
d.LastReason = "apply retry cooldown"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
if now.Sub(d.LastSwitchAt) < bestExitSwitchCooldown {
|
||||
d.LastReason = "cooldown"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
current := findBestExitScore(ordered, d.AppliedExitNodeID)
|
||||
if !bestExitHasMinimumAdvantage(candidate, current) {
|
||||
d.PendingExitNodeID = 0
|
||||
d.PendingCount = 0
|
||||
d.LastReason = "insufficient advantage"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
if d.PendingExitNodeID != candidate.ExitNodeID {
|
||||
d.PendingExitNodeID = candidate.ExitNodeID
|
||||
d.PendingCount = 1
|
||||
d.LastReason = "candidate pending confirmation"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
d.PendingCount++
|
||||
if d.PendingCount < bestExitConfirmationRounds {
|
||||
d.LastReason = "candidate pending confirmation"
|
||||
return bestExitDecisionResult(false, 0, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
d.LastReason = "switch confirmed"
|
||||
return bestExitDecisionResult(true, candidate.ExitNodeID, d.LastReason, ordered)
|
||||
}
|
||||
|
||||
func findBestExitScore(scores []bestExitCandidateScore, exitNodeID int64) bestExitCandidateScore {
|
||||
for _, score := range scores {
|
||||
if score.ExitNodeID == exitNodeID {
|
||||
return score
|
||||
}
|
||||
}
|
||||
return failedBestExitCandidate(0, chainNodeRecord{NodeID: exitNodeID}, "current exit has no successful score")
|
||||
}
|
||||
|
||||
func (m *bestExitManager) decisionLocked(key bestExitOwnerKey) *bestExitDecision {
|
||||
if d := m.decisions[key]; d != nil {
|
||||
return d
|
||||
}
|
||||
d := &bestExitDecision{}
|
||||
m.decisions[key] = d
|
||||
return d
|
||||
}
|
||||
|
||||
func (m *bestExitManager) orderTargets(key bestExitOwnerKey, targets []tunnelRuntimeNode) []tunnelRuntimeNode {
|
||||
out := append([]tunnelRuntimeNode(nil), targets...)
|
||||
if m == nil || len(out) <= 1 {
|
||||
return out
|
||||
}
|
||||
m.mu.Lock()
|
||||
applied := int64(0)
|
||||
if d := m.decisions[key]; d != nil {
|
||||
applied = d.AppliedExitNodeID
|
||||
}
|
||||
m.mu.Unlock()
|
||||
if applied <= 0 {
|
||||
return out
|
||||
}
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
if out[i].NodeID == applied {
|
||||
return true
|
||||
}
|
||||
if out[j].NodeID == applied {
|
||||
return false
|
||||
}
|
||||
return false
|
||||
})
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,248 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
bestExitDisplayStatusApplied = "applied"
|
||||
bestExitDisplayStatusWaiting = "waiting"
|
||||
bestExitDisplaySummaryMulti = "多个出口"
|
||||
bestExitDisplaySummaryWait = "等待探测"
|
||||
bestExitUnknownExitName = "未知出口"
|
||||
bestExitUnknownEntryName = "未知入口"
|
||||
bestExitUnknownChainName = "未知中转"
|
||||
)
|
||||
|
||||
type bestExitDecisionSnapshot struct {
|
||||
AppliedExitNodeID int64
|
||||
UpdatedAt int64
|
||||
Reason string
|
||||
Scores []bestExitCandidateScore
|
||||
}
|
||||
|
||||
type bestExitDisplayState struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Summary string `json:"summary"`
|
||||
Status string `json:"status"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
Items []bestExitDisplayItem `json:"items"`
|
||||
}
|
||||
|
||||
type bestExitDisplayItem struct {
|
||||
OwnerNodeID int64 `json:"ownerNodeId"`
|
||||
OwnerNodeName string `json:"ownerNodeName"`
|
||||
OwnerRole string `json:"ownerRole"`
|
||||
ExitNodeID int64 `json:"exitNodeId,omitempty"`
|
||||
ExitNodeName string `json:"exitNodeName"`
|
||||
UpdatedAt int64 `json:"updatedAt,omitempty"`
|
||||
Reason string `json:"reason,omitempty"`
|
||||
}
|
||||
|
||||
type bestExitNodeNameLookup func(nodeID int64) (string, bool)
|
||||
|
||||
func (m *bestExitManager) snapshot(key bestExitOwnerKey) (bestExitDecisionSnapshot, bool) {
|
||||
if m == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
d := m.decisions[key]
|
||||
if d == nil {
|
||||
return bestExitDecisionSnapshot{}, false
|
||||
}
|
||||
updatedAt := int64(0)
|
||||
if !d.LastSwitchAt.IsZero() {
|
||||
updatedAt = d.LastSwitchAt.UnixMilli()
|
||||
}
|
||||
return bestExitDecisionSnapshot{
|
||||
AppliedExitNodeID: d.AppliedExitNodeID,
|
||||
UpdatedAt: updatedAt,
|
||||
Reason: d.LastReason,
|
||||
Scores: cloneBestExitScores(d.Scores),
|
||||
}, true
|
||||
}
|
||||
|
||||
func (h *Handler) attachBestExitStates(items []map[string]interface{}) {
|
||||
if h == nil || len(items) == 0 {
|
||||
return
|
||||
}
|
||||
lookup := h.bestExitNodeNameLookup()
|
||||
for _, item := range items {
|
||||
state, ok := buildBestExitDisplayState(item, h.bestExit, lookup)
|
||||
if !ok {
|
||||
delete(item, "bestExitState")
|
||||
continue
|
||||
}
|
||||
item["bestExitState"] = state
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) bestExitNodeNameLookup() bestExitNodeNameLookup {
|
||||
cache := map[int64]string{}
|
||||
return func(nodeID int64) (string, bool) {
|
||||
if nodeID <= 0 || h == nil {
|
||||
return "", false
|
||||
}
|
||||
if name, ok := cache[nodeID]; ok {
|
||||
return name, name != ""
|
||||
}
|
||||
node, err := h.getNodeRecord(nodeID)
|
||||
if err != nil || node == nil {
|
||||
cache[nodeID] = ""
|
||||
return "", false
|
||||
}
|
||||
name := strings.TrimSpace(node.Name)
|
||||
cache[nodeID] = name
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
|
||||
func buildBestExitDisplayState(tunnel map[string]interface{}, manager *bestExitManager, lookup bestExitNodeNameLookup) (*bestExitDisplayState, bool) {
|
||||
if tunnel == nil {
|
||||
return nil, false
|
||||
}
|
||||
tunnelID := asInt64(tunnel["id"], 0)
|
||||
outNodes := bestExitDisplayMapSlice(tunnel["outNodeId"])
|
||||
if tunnelID <= 0 || len(outNodes) <= 1 {
|
||||
return nil, false
|
||||
}
|
||||
if !isBestTunnelStrategy(asString(outNodes[0]["strategy"])) {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
owners, ownerRole := bestExitDisplayOwners(tunnel)
|
||||
state := &bestExitDisplayState{
|
||||
Enabled: true,
|
||||
Summary: bestExitDisplaySummaryWait,
|
||||
Status: bestExitDisplayStatusWaiting,
|
||||
Items: make([]bestExitDisplayItem, 0, len(owners)),
|
||||
}
|
||||
|
||||
exitsByID := map[int64]map[string]interface{}{}
|
||||
for _, exit := range outNodes {
|
||||
if id := asInt64(exit["nodeId"], 0); id > 0 {
|
||||
exitsByID[id] = exit
|
||||
}
|
||||
}
|
||||
appliedExitIDs := map[int64]string{}
|
||||
appliedCount := 0
|
||||
latestUpdatedAt := int64(0)
|
||||
latestReason := ""
|
||||
for _, owner := range owners {
|
||||
ownerNodeID := asInt64(owner["nodeId"], 0)
|
||||
if ownerNodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
item := bestExitDisplayItem{
|
||||
OwnerNodeID: ownerNodeID,
|
||||
OwnerNodeName: bestExitDisplayNodeName(owner, ownerNodeID, lookup, bestExitUnknownOwnerName(ownerRole)),
|
||||
OwnerRole: ownerRole,
|
||||
ExitNodeName: bestExitDisplaySummaryWait,
|
||||
Reason: bestExitDisplayStatusWaiting,
|
||||
}
|
||||
if snapshot, ok := manager.snapshot(bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: ownerNodeID}); ok && snapshot.AppliedExitNodeID > 0 {
|
||||
exit, ok := exitsByID[snapshot.AppliedExitNodeID]
|
||||
if !ok {
|
||||
state.Items = append(state.Items, item)
|
||||
continue
|
||||
}
|
||||
item.ExitNodeID = snapshot.AppliedExitNodeID
|
||||
item.ExitNodeName = bestExitDisplayNodeName(exit, snapshot.AppliedExitNodeID, lookup, bestExitUnknownExitName)
|
||||
item.UpdatedAt = snapshot.UpdatedAt
|
||||
item.Reason = snapshot.Reason
|
||||
appliedExitIDs[item.ExitNodeID] = item.ExitNodeName
|
||||
appliedCount++
|
||||
if snapshot.UpdatedAt > latestUpdatedAt {
|
||||
latestUpdatedAt = snapshot.UpdatedAt
|
||||
latestReason = snapshot.Reason
|
||||
}
|
||||
}
|
||||
state.Items = append(state.Items, item)
|
||||
}
|
||||
|
||||
if appliedCount == 0 {
|
||||
return state, true
|
||||
}
|
||||
if appliedCount < len(state.Items) {
|
||||
return state, true
|
||||
}
|
||||
state.Status = bestExitDisplayStatusApplied
|
||||
state.UpdatedAt = latestUpdatedAt
|
||||
state.Reason = latestReason
|
||||
if len(appliedExitIDs) == 1 {
|
||||
for _, name := range appliedExitIDs {
|
||||
state.Summary = name
|
||||
}
|
||||
} else {
|
||||
state.Summary = bestExitDisplaySummaryMulti
|
||||
}
|
||||
return state, true
|
||||
}
|
||||
|
||||
func bestExitDisplayOwners(tunnel map[string]interface{}) ([]map[string]interface{}, string) {
|
||||
chainGroups := bestExitDisplayChainGroups(tunnel["chainNodes"])
|
||||
if len(chainGroups) > 0 {
|
||||
return chainGroups[len(chainGroups)-1], "chain"
|
||||
}
|
||||
return bestExitDisplayMapSlice(tunnel["inNodeId"]), "entry"
|
||||
}
|
||||
|
||||
func bestExitDisplayMapSlice(v interface{}) []map[string]interface{} {
|
||||
switch arr := v.(type) {
|
||||
case []map[string]interface{}:
|
||||
return arr
|
||||
case []interface{}:
|
||||
out := make([]map[string]interface{}, 0, len(arr))
|
||||
for _, item := range arr {
|
||||
if m, ok := item.(map[string]interface{}); ok {
|
||||
out = append(out, m)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayChainGroups(v interface{}) [][]map[string]interface{} {
|
||||
switch groups := v.(type) {
|
||||
case [][]map[string]interface{}:
|
||||
return groups
|
||||
case []interface{}:
|
||||
out := make([][]map[string]interface{}, 0, len(groups))
|
||||
for _, group := range groups {
|
||||
items := bestExitDisplayMapSlice(group)
|
||||
if len(items) > 0 {
|
||||
out = append(out, items)
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func bestExitDisplayNodeName(source map[string]interface{}, nodeID int64, lookup bestExitNodeNameLookup, fallback string) string {
|
||||
if source != nil {
|
||||
for _, key := range []string{"nodeName", "name"} {
|
||||
if name := strings.TrimSpace(asString(source[key])); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
if lookup != nil {
|
||||
if name, ok := lookup(nodeID); ok && strings.TrimSpace(name) != "" {
|
||||
return strings.TrimSpace(name)
|
||||
}
|
||||
}
|
||||
return fallback
|
||||
}
|
||||
|
||||
func bestExitUnknownOwnerName(role string) string {
|
||||
if role == "chain" {
|
||||
return bestExitUnknownChainName
|
||||
}
|
||||
return bestExitUnknownEntryName
|
||||
}
|
||||
@@ -0,0 +1,383 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBestExitDecisionSnapshotIsDefensiveCopy(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
score := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30, NodeName: "exit-a"}, 10, 0, 20, 0)
|
||||
|
||||
m.observeScores(key, []bestExitCandidateScore{score}, now)
|
||||
snapshot, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected snapshot")
|
||||
}
|
||||
if snapshot.AppliedExitNodeID != 30 || snapshot.UpdatedAt != now.UnixMilli() {
|
||||
t.Fatalf("unexpected snapshot: %+v", snapshot)
|
||||
}
|
||||
if len(snapshot.Scores) != 1 {
|
||||
t.Fatalf("expected one score in snapshot, got %+v", snapshot.Scores)
|
||||
}
|
||||
snapshot.Scores[0].ExitNodeID = 99
|
||||
|
||||
again, ok := m.snapshot(key)
|
||||
if !ok {
|
||||
t.Fatalf("expected second snapshot")
|
||||
}
|
||||
if again.Scores[0].ExitNodeID != 30 {
|
||||
t.Fatalf("snapshot score mutation leaked into manager state: %+v", again.Scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForDirectMultiEntryOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(100, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 11}, 31, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
{"nodeId": int64(11)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 11: "入口 B", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if !state.Enabled || state.Summary != "多个出口" || state.Status != "applied" {
|
||||
t.Fatalf("unexpected state summary: %+v", state)
|
||||
}
|
||||
if state.UpdatedAt != now.Add(time.Second).UnixMilli() {
|
||||
t.Fatalf("expected latest updatedAt, got %d", state.UpdatedAt)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "entry" || state.Items[0].OwnerNodeName != "入口 A" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "entry" || state.Items[1].OwnerNodeName != "入口 B" || state.Items[1].ExitNodeName != "日本节点" {
|
||||
t.Fatalf("unexpected second item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateForFinalChainHopOwners(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(200, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 20}, 30, now)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 88, OwnerNodeID: 21}, 30, now.Add(time.Second))
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(88),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{
|
||||
{{"nodeId": int64(15), "inx": int64(0)}},
|
||||
{{"nodeId": int64(20), "inx": int64(1)}, {"nodeId": int64(21), "inx": int64(1)}},
|
||||
},
|
||||
}
|
||||
names := map[int64]string{20: "中转 M1", 21: "中转 M2", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Summary != "香港节点" || state.Status != "applied" {
|
||||
t.Fatalf("expected single-exit summary, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two final-hop owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].OwnerRole != "chain" || state.Items[0].OwnerNodeName != "中转 M1" || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected first chain owner item: %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].OwnerRole != "chain" || state.Items[1].OwnerNodeName != "中转 M2" || state.Items[1].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("unexpected second chain owner item: %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateWaitingWhenNoAppliedDecisionExists(t *testing.T) {
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, newBestExitManager(), testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected waiting best exit state")
|
||||
}
|
||||
if state.Summary != "等待探测" || state.Status != "waiting" {
|
||||
t.Fatalf("expected waiting state, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 1 || state.Items[0].ExitNodeID != 0 || state.Items[0].ExitNodeName != "等待探测" {
|
||||
t.Fatalf("unexpected waiting item: %+v", state.Items)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateKeepsTopLevelWaitingWhenSomeOwnersPending(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(400, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, now)
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
{"nodeId": int64(11)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 11: "入口 B", 30: "香港节点", 31: "日本节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Status != bestExitDisplayStatusWaiting || state.Summary != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected top-level waiting for partial owner state, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 2 {
|
||||
t.Fatalf("expected two owner items, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].ExitNodeID != 30 || state.Items[0].ExitNodeName != "香港节点" {
|
||||
t.Fatalf("expected first owner applied details to remain visible, got %+v", state.Items[0])
|
||||
}
|
||||
if state.Items[1].ExitNodeID != 0 || state.Items[1].ExitNodeName != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected second owner waiting details, got %+v", state.Items[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateIgnoresAppliedExitRemovedFromTunnel(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
now := time.Unix(500, 0)
|
||||
m.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 99, now)
|
||||
|
||||
tunnel := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(10)},
|
||||
},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
{"nodeId": int64(31), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
"chainNodes": [][]map[string]interface{}{},
|
||||
}
|
||||
names := map[int64]string{10: "入口 A", 30: "香港节点", 31: "日本节点", 99: "已删除节点"}
|
||||
|
||||
state, ok := buildBestExitDisplayState(tunnel, m, testBestExitNameLookup(names))
|
||||
if !ok {
|
||||
t.Fatalf("expected best exit state")
|
||||
}
|
||||
if state.Status != bestExitDisplayStatusWaiting || state.Summary != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected waiting state for stale applied exit, got %+v", state)
|
||||
}
|
||||
if len(state.Items) != 1 {
|
||||
t.Fatalf("expected one item, got %+v", state.Items)
|
||||
}
|
||||
if state.Items[0].ExitNodeID != 0 || state.Items[0].ExitNodeName != bestExitDisplaySummaryWait {
|
||||
t.Fatalf("expected stale exit to be ignored, got %+v", state.Items[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildBestExitDisplayStateSkipsNonBestAndSingleExitTunnels(t *testing.T) {
|
||||
nonBest := map[string]interface{}{
|
||||
"id": int64(77),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": "round"},
|
||||
{"nodeId": int64(31), "strategy": "round"},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(nonBest, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected non-best tunnel to skip state, got %+v", state)
|
||||
}
|
||||
|
||||
singleExit := map[string]interface{}{
|
||||
"id": int64(78),
|
||||
"inNodeId": []map[string]interface{}{{"nodeId": int64(10)}},
|
||||
"outNodeId": []map[string]interface{}{
|
||||
{"nodeId": int64(30), "strategy": tunnelStrategyBest},
|
||||
},
|
||||
}
|
||||
if state, ok := buildBestExitDisplayState(singleExit, newBestExitManager(), testBestExitNameLookup(nil)); ok || state != nil {
|
||||
t.Fatalf("expected single-exit tunnel to skip state, got %+v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelListAttachesBestExitStateOnlyForEligibleTunnels(t *testing.T) {
|
||||
h := setupBestExitTunnelHandler(t)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil)
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelList(res, req)
|
||||
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Data []map[string]any `json:"data"`
|
||||
}
|
||||
decodeBestExitTunnelResponse(t, res, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success response, got code %d", payload.Code)
|
||||
}
|
||||
|
||||
bestTunnel := findTunnelResponseItem(t, payload.Data, 77)
|
||||
if _, ok := bestTunnel["bestExitState"]; !ok {
|
||||
t.Fatalf("expected eligible best multi-exit tunnel to include bestExitState: %+v", bestTunnel)
|
||||
}
|
||||
|
||||
singleExitTunnel := findTunnelResponseItem(t, payload.Data, 78)
|
||||
if _, ok := singleExitTunnel["bestExitState"]; ok {
|
||||
t.Fatalf("expected single-exit tunnel to omit bestExitState: %+v", singleExitTunnel)
|
||||
}
|
||||
|
||||
nonBestTunnel := findTunnelResponseItem(t, payload.Data, 79)
|
||||
if _, ok := nonBestTunnel["bestExitState"]; ok {
|
||||
t.Fatalf("expected non-best tunnel to omit bestExitState: %+v", nonBestTunnel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelGetAttachesBestExitStateToSelectedTunnel(t *testing.T) {
|
||||
h := setupBestExitTunnelHandler(t)
|
||||
|
||||
body := bytes.NewReader([]byte(`{"id":77}`))
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/get", body)
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelGet(res, req)
|
||||
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Data map[string]any `json:"data"`
|
||||
}
|
||||
decodeBestExitTunnelResponse(t, res, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success response, got code %d", payload.Code)
|
||||
}
|
||||
if _, ok := payload.Data["bestExitState"]; !ok {
|
||||
t.Fatalf("expected selected best multi-exit tunnel to include bestExitState: %+v", payload.Data)
|
||||
}
|
||||
}
|
||||
|
||||
func setupBestExitTunnelHandler(t *testing.T) *Handler {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
h := New(r, "secret")
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
insertNode := func(id int64, name string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, id, name, name+"-secret", "10.0.0.1", "10.0.0.1", "", "30000-30010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
insertNode(10, "entry-a")
|
||||
insertNode(30, "exit-a")
|
||||
insertNode(31, "exit-b")
|
||||
insertNode(32, "exit-c")
|
||||
|
||||
insertTunnel := func(id int64, name string) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference)
|
||||
VALUES(?, ?, 1, 1, 'tls', 1, ?, ?, 1, ?, '')
|
||||
`, id, name, now, now, id).Error; err != nil {
|
||||
t.Fatalf("insert tunnel %s: %v", name, err)
|
||||
}
|
||||
}
|
||||
insertTunnel(77, "best-multi")
|
||||
insertTunnel(78, "best-single")
|
||||
insertTunnel(79, "round-multi")
|
||||
|
||||
insertChain := func(tunnelID int64, chainType string, nodeID int64, strategy string, inx int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, ?, ?, 30001, ?, ?, 'tls')
|
||||
`, tunnelID, chainType, nodeID, strategy, inx).Error; err != nil {
|
||||
t.Fatalf("insert chain tunnel %d/%s/%d: %v", tunnelID, chainType, nodeID, err)
|
||||
}
|
||||
}
|
||||
insertChain(77, "1", 10, "round", 1)
|
||||
insertChain(77, "3", 30, tunnelStrategyBest, 1)
|
||||
insertChain(77, "3", 31, tunnelStrategyBest, 2)
|
||||
insertChain(78, "1", 10, "round", 1)
|
||||
insertChain(78, "3", 30, tunnelStrategyBest, 1)
|
||||
insertChain(79, "1", 10, "round", 1)
|
||||
insertChain(79, "3", 31, "round", 1)
|
||||
insertChain(79, "3", 32, "round", 2)
|
||||
|
||||
h.bestExit.setApplied(bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}, 30, time.UnixMilli(now))
|
||||
return h
|
||||
}
|
||||
|
||||
func decodeBestExitTunnelResponse(t *testing.T, res *httptest.ResponseRecorder, v any) {
|
||||
t.Helper()
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
if err := json.NewDecoder(res.Body).Decode(v); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func findTunnelResponseItem(t *testing.T, items []map[string]any, id float64) map[string]any {
|
||||
t.Helper()
|
||||
for _, item := range items {
|
||||
if item["id"] == id {
|
||||
return item
|
||||
}
|
||||
}
|
||||
t.Fatalf("tunnel %.0f not found in response: %+v", id, items)
|
||||
return nil
|
||||
}
|
||||
|
||||
func testBestExitNameLookup(names map[int64]string) bestExitNodeNameLookup {
|
||||
return func(nodeID int64) (string, bool) {
|
||||
name := names[nodeID]
|
||||
return name, name != ""
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,451 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
var errBestExitProbeForTest = errors.New("probe failed")
|
||||
|
||||
func TestBestExitScoreCombinesLatencyAndLoss(t *testing.T) {
|
||||
exit := chainNodeRecord{NodeID: 30, NodeName: "exit-a"}
|
||||
score := scoreBestExitCandidate(10, exit, 25, 2, 80, 3)
|
||||
|
||||
if !score.Success {
|
||||
t.Fatalf("expected successful score")
|
||||
}
|
||||
if score.OwnerNodeID != 10 || score.ExitNodeID != 30 {
|
||||
t.Fatalf("unexpected owner/exit ids: %+v", score)
|
||||
}
|
||||
if score.TotalLatency != 105 {
|
||||
t.Fatalf("expected total latency 105, got %v", score.TotalLatency)
|
||||
}
|
||||
if score.TotalLoss < 4.9 || score.TotalLoss > 5.0 {
|
||||
t.Fatalf("expected combined loss about 4.94, got %v", score.TotalLoss)
|
||||
}
|
||||
if score.Score < 599 || score.Score > 600 {
|
||||
t.Fatalf("expected score about 599, got %v", score.Score)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitScorePenalizesLoss(t *testing.T) {
|
||||
stable := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 80, 0, 80, 0)
|
||||
lowLatencyLossy := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 10, 5, 10, 5)
|
||||
|
||||
if !bestExitScoreLess(stable, lowLatencyLossy) {
|
||||
t.Fatalf("expected stable exit to beat low-latency lossy exit: stable=%+v lossy=%+v", stable, lowLatencyLossy)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitFailedCandidateSortsLast(t *testing.T) {
|
||||
failed := failedBestExitCandidate(10, chainNodeRecord{NodeID: 30}, "dial timeout")
|
||||
good := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 100, 0, 100, 0)
|
||||
|
||||
scores := []bestExitCandidateScore{failed, good}
|
||||
sortBestExitScores(scores)
|
||||
|
||||
if scores[0].ExitNodeID != 31 || scores[1].ExitNodeID != 30 {
|
||||
t.Fatalf("expected good score first and failed score last, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitInitialObservationAppliesWithoutSwitch(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate}, now)
|
||||
if decision.Switch {
|
||||
t.Fatalf("initial observation should not return switch: %+v", decision)
|
||||
}
|
||||
if m.decisions[key].AppliedExitNodeID != 31 {
|
||||
t.Fatalf("expected applied exit 31, got %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionRequiresMinimumAdvantage(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 90, 0, 90, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds+1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate below minimum advantage should not switch after repeated observations: %+v", decision)
|
||||
}
|
||||
}
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("expected applied exit to remain 30, got %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionSwitchesWithMinimumAdvantage(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds-1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate should wait for confirmations before switching: %+v", decision)
|
||||
}
|
||||
}
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds-1)*time.Second))
|
||||
if !decision.Switch || decision.ExitNodeID != 31 {
|
||||
t.Fatalf("candidate with enough advantage should switch after confirmations: %+v", decision)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitConfirmedSwitchDoesNotMarkAppliedUntilSetApplied(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds-1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate should wait for confirmations before switching: %+v", decision)
|
||||
}
|
||||
}
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds-1)*time.Second))
|
||||
if !decision.Switch || decision.ExitNodeID != 31 {
|
||||
t.Fatalf("candidate with enough advantage should switch after confirmations: %+v", decision)
|
||||
}
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("confirmed switch should not mark applied before runtime update: %+v", m.decisions[key])
|
||||
}
|
||||
|
||||
m.setApplied(key, decision.ExitNodeID, now.Add(time.Second))
|
||||
if m.decisions[key].AppliedExitNodeID != 31 {
|
||||
t.Fatalf("setApplied should commit confirmed switch: %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitApplyFailureStartsRetryCooldownWithoutChangingAppliedExit(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
for i := 0; i < bestExitConfirmationRounds-1; i++ {
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Duration(i)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("candidate should wait for confirmations before switching: %+v", decision)
|
||||
}
|
||||
}
|
||||
confirmed := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds-1)*time.Second))
|
||||
if !confirmed.Switch || confirmed.ExitNodeID != 31 {
|
||||
t.Fatalf("expected confirmed switch before apply failure: %+v", confirmed)
|
||||
}
|
||||
|
||||
m.recordApplyFailure(key, confirmed.ExitNodeID, now.Add(bestExitConfirmationRounds*time.Second))
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("apply failure should leave applied exit unchanged: %+v", m.decisions[key])
|
||||
}
|
||||
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add((bestExitConfirmationRounds+1)*time.Second))
|
||||
if decision.Switch {
|
||||
t.Fatalf("apply retry cooldown should suppress immediate retry: %+v", decision)
|
||||
}
|
||||
if decision.Reason != "apply retry cooldown" {
|
||||
t.Fatalf("expected apply retry cooldown reason, got %q", decision.Reason)
|
||||
}
|
||||
|
||||
retry := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(bestExitConfirmationRounds*time.Second+bestExitApplyRetryCooldown))
|
||||
if !retry.Switch || retry.ExitNodeID != 31 {
|
||||
t.Fatalf("expected retry after apply cooldown: %+v", retry)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitEnsureAppliedDoesNotOverrideExistingAppliedExit(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
|
||||
m.ensureApplied(key, 30, now)
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("expected initial applied exit 30, got %+v", m.decisions[key])
|
||||
}
|
||||
if !m.decisions[key].LastSwitchAt.Equal(now) {
|
||||
t.Fatalf("expected initial applied timestamp, got %+v", m.decisions[key])
|
||||
}
|
||||
|
||||
m.ensureApplied(key, 31, now.Add(time.Minute))
|
||||
if m.decisions[key].AppliedExitNodeID != 30 {
|
||||
t.Fatalf("ensureApplied should not override existing applied exit: %+v", m.decisions[key])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitRoundPingerCachesByNodeHostAndPort(t *testing.T) {
|
||||
publicCalls := 0
|
||||
ownerCalls := 0
|
||||
pinger := newBestExitRoundPinger(func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
if ip == bestExitPublicTargetHost && port == bestExitPublicTargetPort {
|
||||
publicCalls++
|
||||
return float64(nodeID), 0, nil
|
||||
}
|
||||
ownerCalls++
|
||||
return float64(ownerCalls), 0, nil
|
||||
})
|
||||
|
||||
if lat, _, err := pinger(30, bestExitPublicTargetHost, bestExitPublicTargetPort, diagnosisExecOptions{}); err != nil || lat != 30 {
|
||||
t.Fatalf("unexpected first public ping result lat=%v err=%v", lat, err)
|
||||
}
|
||||
if lat, _, err := pinger(30, bestExitPublicTargetHost, bestExitPublicTargetPort, diagnosisExecOptions{}); err != nil || lat != 30 {
|
||||
t.Fatalf("unexpected cached public ping result lat=%v err=%v", lat, err)
|
||||
}
|
||||
if _, _, err := pinger(31, bestExitPublicTargetHost, bestExitPublicTargetPort, diagnosisExecOptions{}); err != nil {
|
||||
t.Fatalf("unexpected second exit public ping err=%v", err)
|
||||
}
|
||||
if publicCalls != 2 {
|
||||
t.Fatalf("expected public probes cached per exit node, got %d calls", publicCalls)
|
||||
}
|
||||
|
||||
if _, _, err := pinger(10, "10.0.0.30", 30030, diagnosisExecOptions{}); err != nil {
|
||||
t.Fatalf("unexpected owner ping err=%v", err)
|
||||
}
|
||||
if _, _, err := pinger(10, "10.0.0.30", 30030, diagnosisExecOptions{}); err != nil {
|
||||
t.Fatalf("unexpected repeated owner ping err=%v", err)
|
||||
}
|
||||
if ownerCalls != 1 {
|
||||
t.Fatalf("expected owner-to-exit probes cached by target, got %d calls", ownerCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionScoresAreDefensiveCopies(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now)
|
||||
decision.Scores[0].ExitNodeID = 99
|
||||
|
||||
if m.decisions[key].Scores[0].ExitNodeID != 31 {
|
||||
t.Fatalf("decision scores mutation leaked into manager state: %+v", m.decisions[key].Scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitDecisionRequiresConfirmationsAndCooldown(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
now := time.Unix(100, 0)
|
||||
current := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 100, 0, 100, 0)
|
||||
candidate := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 31}, 40, 0, 60, 0)
|
||||
|
||||
m.setApplied(key, 30, now.Add(-time.Minute))
|
||||
|
||||
if decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now); decision.Switch {
|
||||
t.Fatalf("first observation should not switch: %+v", decision)
|
||||
}
|
||||
if decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(time.Second)); decision.Switch {
|
||||
t.Fatalf("second observation should not switch: %+v", decision)
|
||||
}
|
||||
decision := m.observeScores(key, []bestExitCandidateScore{candidate, current}, now.Add(2*time.Second))
|
||||
if !decision.Switch || decision.ExitNodeID != 31 {
|
||||
t.Fatalf("third confirmed observation should switch to 31: %+v", decision)
|
||||
}
|
||||
|
||||
betterAgain := scoreBestExitCandidate(10, chainNodeRecord{NodeID: 30}, 20, 0, 20, 0)
|
||||
if decision := m.observeScores(key, []bestExitCandidateScore{betterAgain, candidate}, now.Add(3*time.Second)); decision.Switch {
|
||||
t.Fatalf("cooldown should block immediate switch back: %+v", decision)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBestExitOrderingUsesAppliedDecision(t *testing.T) {
|
||||
m := newBestExitManager()
|
||||
key := bestExitOwnerKey{TunnelID: 7, OwnerNodeID: 10}
|
||||
m.setApplied(key, 31, time.Unix(100, 0))
|
||||
targets := []tunnelRuntimeNode{
|
||||
{NodeID: 30, Strategy: tunnelStrategyBest},
|
||||
{NodeID: 31, Strategy: tunnelStrategyBest},
|
||||
{NodeID: 32, Strategy: tunnelStrategyBest},
|
||||
}
|
||||
|
||||
ordered := m.orderTargets(key, targets)
|
||||
if ordered[0].NodeID != 31 || ordered[1].NodeID != 30 || ordered[2].NodeID != 32 {
|
||||
t.Fatalf("unexpected order: %+v", ordered)
|
||||
}
|
||||
if targets[0].NodeID != 30 {
|
||||
t.Fatalf("orderTargets mutated input: %+v", targets)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTunnelChainConfigMapsBestStrategyToFIFO(t *testing.T) {
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
31: {ID: 31, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
|
||||
}
|
||||
targets := []tunnelRuntimeNode{
|
||||
{NodeID: 30, Port: 30030, Protocol: "tls", Strategy: tunnelStrategyBest, ChainType: 3},
|
||||
{NodeID: 31, Port: 30031, Protocol: "tls", Strategy: tunnelStrategyBest, ChainType: 3},
|
||||
}
|
||||
|
||||
chainData, err := buildTunnelChainConfig(77, 10, targets, nodes, "")
|
||||
if err != nil {
|
||||
t.Fatalf("build chain: %v", err)
|
||||
}
|
||||
hops := chainData["hops"].([]map[string]interface{})
|
||||
selector := hops[0]["selector"].(map[string]interface{})
|
||||
if selector["strategy"] != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best to render as fifo, got %v", selector["strategy"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerOrdersBestExitTargetsForOwner(t *testing.T) {
|
||||
h := &Handler{bestExit: newBestExitManager()}
|
||||
key := bestExitOwnerKey{TunnelID: 77, OwnerNodeID: 10}
|
||||
h.bestExit.setApplied(key, 31, time.Unix(100, 0))
|
||||
targets := []tunnelRuntimeNode{
|
||||
{NodeID: 30, Port: 30030, Strategy: tunnelStrategyBest},
|
||||
{NodeID: 31, Port: 30031, Strategy: tunnelStrategyBest},
|
||||
}
|
||||
|
||||
ordered := h.orderBestExitTargets(77, 10, targets)
|
||||
if ordered[0].NodeID != 31 || ordered[1].NodeID != 30 {
|
||||
t.Fatalf("unexpected ordered targets: %+v", ordered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeStrategyForTargetsMapsBestTargetStrategyToFIFO(t *testing.T) {
|
||||
owner := tunnelRuntimeNode{Strategy: "round"}
|
||||
targets := []tunnelRuntimeNode{{Strategy: tunnelStrategyBest}}
|
||||
|
||||
if got := runtimeStrategyForTargets(owner, targets); got != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best target strategy to map to fifo, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeStrategyForTargetsPreservesNonBestTargetStrategy(t *testing.T) {
|
||||
owner := tunnelRuntimeNode{Strategy: tunnelStrategyBest}
|
||||
targets := []tunnelRuntimeNode{{Strategy: "round"}}
|
||||
|
||||
if got := runtimeStrategyForTargets(owner, targets); got != "round" {
|
||||
t.Fatalf("expected target strategy round to remain unchanged, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeStrategyForTargetsMapsBestOwnerStrategyWhenTargetsEmpty(t *testing.T) {
|
||||
owner := tunnelRuntimeNode{Strategy: tunnelStrategyBest}
|
||||
|
||||
if got := runtimeStrategyForTargets(owner, nil); got != bestExitRuntimeStrategy {
|
||||
t.Fatalf("expected best owner fallback strategy to map to fifo, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{
|
||||
{NodeID: 30, NodeName: "exit-a", Port: 30030},
|
||||
{NodeID: 31, NodeName: "exit-b", Port: 30031},
|
||||
}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
31: {ID: 31, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
switch {
|
||||
case nodeID == 10 && port == 30030:
|
||||
return 60, 0, nil
|
||||
case nodeID == 10 && port == 30031:
|
||||
return 20, 0, nil
|
||||
case nodeID == 30 && ip == bestExitPublicTargetHost:
|
||||
return 60, 0, nil
|
||||
case nodeID == 31 && ip == bestExitPublicTargetHost:
|
||||
return 20, 0, nil
|
||||
default:
|
||||
t.Fatalf("unexpected ping node=%d ip=%s port=%d", nodeID, ip, port)
|
||||
return 0, 100, nil
|
||||
}
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), pinger)
|
||||
if len(scores) != 2 {
|
||||
t.Fatalf("expected two scores, got %+v", scores)
|
||||
}
|
||||
if scores[0].ExitNodeID != 31 {
|
||||
t.Fatalf("expected exit-b first, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerUsesConfiguredPublicProbeTarget(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry-a"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30001}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Name: "entry-a", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Name: "exit-a", ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
}
|
||||
target := tunnelProbeTarget{Host: "speed.example.com", Port: 8443}
|
||||
var calls []string
|
||||
ping := func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, target, ping)
|
||||
if len(scores) != 1 || !scores[0].Success {
|
||||
t.Fatalf("expected successful score, got %+v", scores)
|
||||
}
|
||||
if !slices.Contains(calls, "30|speed.example.com|8443") {
|
||||
t.Fatalf("expected exit public probe to use configured target, calls=%+v", calls)
|
||||
}
|
||||
for _, call := range calls {
|
||||
if strings.Contains(call, defaultTunnelProbeTargetHost) {
|
||||
t.Fatalf("did not expect default target call when custom target configured: %+v", calls)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerMarksCandidateFailedWhenOwnerToExitFails(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
return 0, 100, errBestExitProbeForTest
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), pinger)
|
||||
if len(scores) != 1 || scores[0].Success {
|
||||
t.Fatalf("expected failed candidate, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerMarksCandidateFailedWhenTargetResolutionFails(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-v6", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Name: "entry", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Name: "exit-v6", ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
t.Fatalf("ping should not be called when target resolution fails: node=%d ip=%s port=%d", nodeID, ip, port)
|
||||
return 0, 100, nil
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "v4", diagnosisExecOptions{}, defaultTunnelProbeTarget(), pinger)
|
||||
if len(scores) != 1 || scores[0].Success {
|
||||
t.Fatalf("expected failed candidate, got %+v", scores)
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type tunnelTrafficDelta struct {
|
||||
@@ -21,75 +22,42 @@ func unixMilliBucketMinute(nowMs int64) int64 {
|
||||
return nowMs - (nowMs % minuteMs)
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowItem, nowMs int64) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if nodeID <= 0 || len(items) == 0 {
|
||||
return
|
||||
func collectFlowUploadForwardIDs(items []flowItem) []int64 {
|
||||
ids := make([]int64, 0, len(items))
|
||||
seen := make(map[int64]struct{}, len(items))
|
||||
for _, item := range items {
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(strings.TrimSpace(item.N))
|
||||
if !ok || forwardID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, exists := seen[forwardID]; exists {
|
||||
continue
|
||||
}
|
||||
seen[forwardID] = struct{}{}
|
||||
ids = append(ids, forwardID)
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
func (h *Handler) recordTunnelMetricsFromForwardBatch(nodeID int64, forwardDeltas map[int64]tunnelTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta, nowMs int64) {
|
||||
if h == nil || h.repo == nil || nodeID <= 0 || len(forwardDeltas) == 0 {
|
||||
return
|
||||
}
|
||||
bucketTs := unixMilliBucketMinute(nowMs)
|
||||
if bucketTs <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
forwardDeltas := make(map[int64]tunnelTrafficDelta)
|
||||
var skippedParse, skippedZero int
|
||||
for _, item := range items {
|
||||
name := strings.TrimSpace(item.N)
|
||||
if name == "" || name == "web_api" {
|
||||
continue
|
||||
}
|
||||
forwardID, _, _, ok := parseFlowServiceIDs(name)
|
||||
if !ok {
|
||||
skippedParse++
|
||||
continue
|
||||
}
|
||||
if item.D == 0 && item.U == 0 {
|
||||
skippedZero++
|
||||
continue
|
||||
}
|
||||
d := forwardDeltas[forwardID]
|
||||
d.bytesIn += item.D
|
||||
d.bytesOut += item.U
|
||||
forwardDeltas[forwardID] = d
|
||||
}
|
||||
if len(forwardDeltas) == 0 {
|
||||
if len(items) > 0 {
|
||||
log.Printf("monitoring debug op=tunnel_metric.no_forward_deltas node_id=%d items=%d skipped_parse=%d skipped_zero=%d", nodeID, len(items), skippedParse, skippedZero)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
forwardIDs := make([]int64, 0, len(forwardDeltas))
|
||||
for id := range forwardDeltas {
|
||||
forwardIDs = append(forwardIDs, id)
|
||||
}
|
||||
|
||||
forwardTunnelMap, err := h.repo.MapForwardIDsToTunnelIDs(forwardIDs)
|
||||
if err != nil {
|
||||
log.Printf("monitoring write skipped op=tunnel_metric.map_forward_to_tunnel node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
if len(forwardTunnelMap) == 0 {
|
||||
log.Printf("monitoring debug op=tunnel_metric.no_tunnel_map node_id=%d forward_ids=%v", nodeID, forwardIDs)
|
||||
return
|
||||
}
|
||||
|
||||
tunnelAgg := make(map[int64]tunnelTrafficDelta)
|
||||
for forwardID, delta := range forwardDeltas {
|
||||
tunnelID := forwardTunnelMap[forwardID]
|
||||
if tunnelID <= 0 {
|
||||
meta, ok := metas[forwardID]
|
||||
if !ok || meta.TunnelID <= 0 {
|
||||
continue
|
||||
}
|
||||
a := tunnelAgg[tunnelID]
|
||||
a.bytesIn += delta.bytesIn
|
||||
a.bytesOut += delta.bytesOut
|
||||
tunnelAgg[tunnelID] = a
|
||||
}
|
||||
if len(tunnelAgg) == 0 {
|
||||
return
|
||||
current := tunnelAgg[meta.TunnelID]
|
||||
current.bytesIn += delta.bytesIn
|
||||
current.bytesOut += delta.bytesOut
|
||||
tunnelAgg[meta.TunnelID] = current
|
||||
}
|
||||
|
||||
metrics := make([]*model.TunnelMetric, 0, len(tunnelAgg))
|
||||
@@ -98,14 +66,11 @@ func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowIte
|
||||
continue
|
||||
}
|
||||
metrics = append(metrics, &model.TunnelMetric{
|
||||
TunnelID: tunnelID,
|
||||
NodeID: nodeID,
|
||||
Timestamp: bucketTs,
|
||||
BytesIn: delta.bytesIn,
|
||||
BytesOut: delta.bytesOut,
|
||||
Connections: 0,
|
||||
Errors: 0,
|
||||
AvgLatencyMs: 0,
|
||||
TunnelID: tunnelID,
|
||||
NodeID: nodeID,
|
||||
Timestamp: bucketTs,
|
||||
BytesIn: delta.bytesIn,
|
||||
BytesOut: delta.bytesOut,
|
||||
})
|
||||
}
|
||||
if len(metrics) == 0 {
|
||||
@@ -114,7 +79,7 @@ func (h *Handler) recordTunnelMetricsFromFlowItems(nodeID int64, items []flowIte
|
||||
|
||||
if err := h.repo.UpsertTunnelMetricBuckets(metrics); err != nil {
|
||||
log.Printf("monitoring write failed op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d err=%v", nodeID, bucketTs, len(metrics), err)
|
||||
} else {
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
return
|
||||
}
|
||||
log.Printf("monitoring ok op=tunnel_metric.upsert_buckets node_id=%d bucket_ts=%d count=%d", nodeID, bucketTs, len(metrics))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTunnelProbeTargetHost = "www.bing.com"
|
||||
defaultTunnelProbeTargetPort = 443
|
||||
)
|
||||
|
||||
type tunnelProbeTarget struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
func defaultTunnelProbeTarget() tunnelProbeTarget {
|
||||
return tunnelProbeTarget{Host: defaultTunnelProbeTargetHost, Port: defaultTunnelProbeTargetPort}
|
||||
}
|
||||
|
||||
func normalizeTunnelProbeTarget(host string, port int) (tunnelProbeTarget, bool, error) {
|
||||
host = strings.TrimSpace(host)
|
||||
if host == "" && port == 0 {
|
||||
return defaultTunnelProbeTarget(), false, nil
|
||||
}
|
||||
if host == "" {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标 Host 不能为空")
|
||||
}
|
||||
if port <= 0 || port > 65535 {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标端口必须是 1-65535")
|
||||
}
|
||||
if strings.Contains(host, "://") || strings.ContainsAny(host, "/?#") || strings.ContainsAny(host, " \t\r\n") || isTunnelProbeTargetSchemeLikeHost(host) {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标 Host 不能包含协议或路径")
|
||||
}
|
||||
if normalized, ok := normalizeTunnelProbeTargetHost(host); ok {
|
||||
host = normalized
|
||||
} else {
|
||||
return tunnelProbeTarget{}, false, errors.New("测试目标 Host 格式无效")
|
||||
}
|
||||
|
||||
return tunnelProbeTarget{Host: host, Port: port}, true, nil
|
||||
}
|
||||
|
||||
func normalizeTunnelProbeTargetHost(host string) (string, bool) {
|
||||
if strings.HasPrefix(host, "[") || strings.HasSuffix(host, "]") {
|
||||
if !strings.HasPrefix(host, "[") || !strings.HasSuffix(host, "]") {
|
||||
return "", false
|
||||
}
|
||||
inner := strings.TrimPrefix(strings.TrimSuffix(host, "]"), "[")
|
||||
addr, err := netip.ParseAddr(inner)
|
||||
if err != nil || !addr.Is6() {
|
||||
return "", false
|
||||
}
|
||||
return inner, true
|
||||
}
|
||||
|
||||
if addr, err := netip.ParseAddr(host); err == nil {
|
||||
return addr.String(), true
|
||||
}
|
||||
if strings.Contains(host, ":") || isTunnelProbeTargetIPv4Like(host) {
|
||||
return "", false
|
||||
}
|
||||
if !isValidTunnelProbeTargetHost(host) {
|
||||
return "", false
|
||||
}
|
||||
return host, true
|
||||
}
|
||||
|
||||
func isValidTunnelProbeTargetHost(host string) bool {
|
||||
if host == "" || len(host) > 253 {
|
||||
return false
|
||||
}
|
||||
for _, label := range strings.Split(host, ".") {
|
||||
if len(label) == 0 || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return false
|
||||
}
|
||||
for _, r := range label {
|
||||
if !isASCIILetter(r) && !isASCIIDigit(r) && r != '-' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func isTunnelProbeTargetIPv4Like(host string) bool {
|
||||
if host == "" {
|
||||
return false
|
||||
}
|
||||
for _, r := range host {
|
||||
if !isASCIIDigit(r) && r != '.' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return strings.Contains(host, ".")
|
||||
}
|
||||
|
||||
func isTunnelProbeTargetSchemeLikeHost(host string) bool {
|
||||
if _, err := netip.ParseAddr(host); err == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
colon := strings.IndexByte(host, ':')
|
||||
if colon <= 0 {
|
||||
return false
|
||||
}
|
||||
for i, r := range host[:colon] {
|
||||
if i == 0 {
|
||||
if !isASCIILetter(r) {
|
||||
return false
|
||||
}
|
||||
continue
|
||||
}
|
||||
if !isASCIILetter(r) && !isASCIIDigit(r) && r != '+' && r != '-' && r != '.' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func isASCIILetter(r rune) bool {
|
||||
return (r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z')
|
||||
}
|
||||
|
||||
func isASCIIDigit(r rune) bool {
|
||||
return r >= '0' && r <= '9'
|
||||
}
|
||||
|
||||
func parseTunnelProbeTargetFromRequest(req map[string]interface{}) (tunnelProbeTarget, bool, error) {
|
||||
if req == nil {
|
||||
return defaultTunnelProbeTarget(), false, nil
|
||||
}
|
||||
rawHost, hasHost := req["probeTargetHost"]
|
||||
rawPort, hasPort := req["probeTargetPort"]
|
||||
if !hasHost && !hasPort {
|
||||
return defaultTunnelProbeTarget(), false, nil
|
||||
}
|
||||
host, err := parseTunnelProbeTargetHostValue(rawHost)
|
||||
if err != nil {
|
||||
return tunnelProbeTarget{}, false, err
|
||||
}
|
||||
port, err := parseTunnelProbeTargetPortValue(rawPort)
|
||||
if err != nil {
|
||||
return tunnelProbeTarget{}, false, err
|
||||
}
|
||||
return normalizeTunnelProbeTarget(host, port)
|
||||
}
|
||||
|
||||
func parseTunnelProbeTargetHostValue(raw interface{}) (string, error) {
|
||||
if raw == nil {
|
||||
return "", nil
|
||||
}
|
||||
host, ok := raw.(string)
|
||||
if !ok {
|
||||
return "", errors.New("测试目标 Host 格式无效")
|
||||
}
|
||||
if host != strings.TrimSpace(host) {
|
||||
return "", errors.New("测试目标 Host 不能包含协议或路径")
|
||||
}
|
||||
return host, nil
|
||||
}
|
||||
|
||||
func parseTunnelProbeTargetPortValue(raw interface{}) (int, error) {
|
||||
if raw == nil {
|
||||
return 0, nil
|
||||
}
|
||||
switch v := raw.(type) {
|
||||
case float64:
|
||||
if v != float64(int64(v)) {
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
return int(v), nil
|
||||
case string:
|
||||
if v == "" {
|
||||
return 0, nil
|
||||
}
|
||||
if v != strings.TrimSpace(v) {
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
port, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
return port, nil
|
||||
case int:
|
||||
return v, nil
|
||||
case int32:
|
||||
return int(v), nil
|
||||
case int64:
|
||||
return int(v), nil
|
||||
default:
|
||||
return 0, errors.New("测试目标端口必须是整数")
|
||||
}
|
||||
}
|
||||
|
||||
func effectiveTunnelProbeTarget(tunnel *model.Tunnel) tunnelProbeTarget {
|
||||
if tunnel == nil {
|
||||
return defaultTunnelProbeTarget()
|
||||
}
|
||||
return effectiveTunnelProbeTargetValues(tunnel.ProbeTargetHost, tunnel.ProbeTargetPort)
|
||||
}
|
||||
|
||||
func effectiveTunnelProbeTargetValues(host string, port int) tunnelProbeTarget {
|
||||
target, configured, err := normalizeTunnelProbeTarget(host, port)
|
||||
if err != nil || !configured {
|
||||
return defaultTunnelProbeTarget()
|
||||
}
|
||||
return target
|
||||
}
|
||||
|
||||
func formatTunnelProbeTarget(target tunnelProbeTarget) string {
|
||||
if addr, err := netip.ParseAddr(target.Host); err == nil && addr.Is6() {
|
||||
return fmt.Sprintf("[%s]:%d", target.Host, target.Port)
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", target.Host, target.Port)
|
||||
}
|
||||
@@ -0,0 +1,305 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestTunnelCreatePersistsProbeTargetAndListReturnsConfiguredValue(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"name":"custom-target",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"speed.example.com",
|
||||
"probeTargetPort":8443
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelCreate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", body))
|
||||
assertProbeTargetSuccess(t, res)
|
||||
|
||||
listRes := httptest.NewRecorder()
|
||||
h.tunnelList(listRes, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/list", nil))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Data []map[string]any `json:"data"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, listRes, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success, got code %d", payload.Code)
|
||||
}
|
||||
item := payload.Data[0]
|
||||
if item["probeTargetHost"] != "speed.example.com" || item["probeTargetPort"] != float64(8443) {
|
||||
t.Fatalf("unexpected probe target in list response: %+v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdatePersistsDefaultProbeTargetAsEmpty(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 77, "existing", "old.example.com", 9443)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":77,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"",
|
||||
"probeTargetPort":0
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
assertProbeTargetSuccess(t, res)
|
||||
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list tunnels: %v", err)
|
||||
}
|
||||
item := findProbeTargetTunnelItem(t, items, 77)
|
||||
if item["probeTargetHost"] != "" || item["probeTargetPort"] != 0 {
|
||||
t.Fatalf("expected default target to round-trip as empty/0, got %+v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateWithoutProbeTargetFieldsPreservesExistingTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 79, "existing", "old.example.com", 9443)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":79,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}]
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
assertProbeTargetSuccess(t, res)
|
||||
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list tunnels: %v", err)
|
||||
}
|
||||
item := findProbeTargetTunnelItem(t, items, 79)
|
||||
if item["probeTargetHost"] != "old.example.com" || item["probeTargetPort"] != 9443 {
|
||||
t.Fatalf("expected omitted probe target fields to preserve existing target, got %+v", item)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateRejectsInvalidProbeTargetWithoutClearingExistingTarget(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
probeFields string
|
||||
}{
|
||||
{name: "non numeric port", probeFields: `,"probeTargetPort":"abc"`},
|
||||
{name: "fractional port", probeFields: `,"probeTargetPort":443.5`},
|
||||
{name: "whitespace host", probeFields: `,"probeTargetHost":" "`},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 80, "existing", "old.example.com", 9443)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":80,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}]
|
||||
` + tt.probeFields + `}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code == 0 || payload.Msg == "" {
|
||||
t.Fatalf("expected validation failure, got %+v", payload)
|
||||
}
|
||||
|
||||
items, err := h.repo.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list tunnels: %v", err)
|
||||
}
|
||||
item := findProbeTargetTunnelItem(t, items, 80)
|
||||
if item["probeTargetHost"] != "old.example.com" || item["probeTargetPort"] != 9443 {
|
||||
t.Fatalf("expected invalid probe target to preserve existing target, got %+v", item)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelCreateRejectsInvalidProbeTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"name":"bad-target",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"https://example.com",
|
||||
"probeTargetPort":443
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelCreate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/create", body))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code == 0 || payload.Msg == "" {
|
||||
t.Fatalf("expected validation failure, got %+v", payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateInvalidProbeTargetDoesNotCleanFederationBindings(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 88, "existing", "old.example.com", 9443)
|
||||
seedProbeTargetFederationBinding(t, h, 88)
|
||||
body := bytes.NewReader([]byte(`{
|
||||
"id":88,
|
||||
"name":"existing",
|
||||
"type":1,
|
||||
"flow":1,
|
||||
"trafficRatio":1,
|
||||
"status":1,
|
||||
"inNodeId":[{"nodeId":10,"protocol":"tls"}],
|
||||
"probeTargetHost":"https://example.com",
|
||||
"probeTargetPort":443
|
||||
}`))
|
||||
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelUpdate(res, httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/update", body))
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code == 0 || payload.Msg == "" {
|
||||
t.Fatalf("expected validation failure, got %+v", payload)
|
||||
}
|
||||
|
||||
bindings, err := h.repo.ListActiveFederationTunnelBindingsByTunnel(88)
|
||||
if err != nil {
|
||||
t.Fatalf("list federation bindings: %v", err)
|
||||
}
|
||||
if len(bindings) != 1 {
|
||||
t.Fatalf("expected federation binding to remain after invalid update, got %d", len(bindings))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelDiagnosisUsesConfiguredProbeTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 90, "diagnosis-target", "speed.example.com", 8443)
|
||||
|
||||
_, _, workItems, err := h.prepareTunnelDiagnosis(90)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare tunnel diagnosis: %v", err)
|
||||
}
|
||||
if len(workItems) != 1 {
|
||||
t.Fatalf("expected one diagnosis item, got %d", len(workItems))
|
||||
}
|
||||
if workItems[0].targetIP != "speed.example.com" || workItems[0].targetPort != 8443 {
|
||||
t.Fatalf("expected custom diagnosis target speed.example.com:8443, got %s:%d", workItems[0].targetIP, workItems[0].targetPort)
|
||||
}
|
||||
}
|
||||
|
||||
func setupProbeTargetTunnelHandler(t *testing.T) *Handler {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
h := New(r, "secret")
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(10, 'entry-a', 'entry-secret', '10.0.0.1', '10.0.0.1', '', '30000-30010', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
func seedProbeTargetTunnel(t *testing.T, h *Handler, id int64, name string, host string, port int) {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference, probe_target_host, probe_target_port)
|
||||
VALUES(?, ?, 1, 1, 'tls', 1, ?, ?, 1, ?, '', ?, ?)
|
||||
`, id, name, now, now, id, host, port).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, '1', 10, 30001, 'round', 1, 'tls')
|
||||
`, id).Error; err != nil {
|
||||
t.Fatalf("insert chain: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func seedProbeTargetFederationBinding(t *testing.T, h *Handler, tunnelID int64) {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time)
|
||||
VALUES(?, 10, 1, 0, 'http://peer.example', ?, 'remote-binding', 30001, 1, ?, ?)
|
||||
`, tunnelID, "probe-target-test-binding", now, now).Error; err != nil {
|
||||
t.Fatalf("insert federation binding: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertProbeTargetSuccess(t *testing.T, res *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
decodeProbeTargetResponse(t, res, &payload)
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success, got %+v", payload)
|
||||
}
|
||||
}
|
||||
|
||||
func decodeProbeTargetResponse(t *testing.T, res *httptest.ResponseRecorder, v any) {
|
||||
t.Helper()
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
if err := json.NewDecoder(res.Body).Decode(v); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func findProbeTargetTunnelItem(t *testing.T, items []map[string]interface{}, id int64) map[string]interface{} {
|
||||
t.Helper()
|
||||
for _, item := range items {
|
||||
if asInt64(item["id"], 0) == id {
|
||||
return item
|
||||
}
|
||||
}
|
||||
t.Fatalf("tunnel %d not found: %+v", id, items)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package handler
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestNormalizeTunnelProbeTargetDefaultsWhenEmpty(t *testing.T) {
|
||||
target, configured, err := normalizeTunnelProbeTarget("", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if configured {
|
||||
t.Fatalf("expected empty input to be default, not configured")
|
||||
}
|
||||
if target.Host != defaultTunnelProbeTargetHost || target.Port != defaultTunnelProbeTargetPort {
|
||||
t.Fatalf("unexpected default target: %+v", target)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetAcceptsHostPortAndIPv6(t *testing.T) {
|
||||
target, configured, err := normalizeTunnelProbeTarget(" [2001:db8::1] ", 8443)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !configured {
|
||||
t.Fatalf("expected explicit target")
|
||||
}
|
||||
if target.Host != "2001:db8::1" || target.Port != 8443 {
|
||||
t.Fatalf("unexpected normalized target: %+v", target)
|
||||
}
|
||||
if got := formatTunnelProbeTarget(target); got != "[2001:db8::1]:8443" {
|
||||
t.Fatalf("unexpected formatted target: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetRejectsPartialAndInvalidInputs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
host string
|
||||
port int
|
||||
}{
|
||||
{name: "missing host", host: "", port: 443},
|
||||
{name: "missing port", host: "example.com", port: 0},
|
||||
{name: "port too high", host: "example.com", port: 70000},
|
||||
{name: "scheme", host: "https://example.com", port: 443},
|
||||
{name: "path", host: "example.com/ping", port: 443},
|
||||
{name: "space", host: "example .com", port: 443},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if _, _, err := normalizeTunnelProbeTarget(tt.host, tt.port); err == nil {
|
||||
t.Fatalf("expected validation error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetRejectsSchemePrefixButAllowsIPv6(t *testing.T) {
|
||||
for _, host := range []string{"https:example.com", "mailto:ops@example.com"} {
|
||||
if _, _, err := normalizeTunnelProbeTarget(host, 443); err == nil {
|
||||
t.Fatalf("expected scheme-like host %q to be rejected", host)
|
||||
}
|
||||
}
|
||||
|
||||
for _, host := range []string{"2001:db8::1", "[2001:db8::1]"} {
|
||||
target, configured, err := normalizeTunnelProbeTarget(host, 443)
|
||||
if err != nil {
|
||||
t.Fatalf("expected IPv6 host %q to be accepted: %v", host, err)
|
||||
}
|
||||
if !configured || target.Host != "2001:db8::1" {
|
||||
t.Fatalf("unexpected IPv6 normalization for %q: %+v configured=%v", host, target, configured)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelProbeTargetValidatesHostShape(t *testing.T) {
|
||||
validHosts := []string{
|
||||
"example.com",
|
||||
"localhost",
|
||||
"api-1.example.co.uk",
|
||||
"192.0.2.10",
|
||||
"2001:db8::1",
|
||||
"[2001:db8::1]",
|
||||
}
|
||||
for _, host := range validHosts {
|
||||
if _, _, err := normalizeTunnelProbeTarget(host, 443); err != nil {
|
||||
t.Fatalf("expected valid host %q: %v", host, err)
|
||||
}
|
||||
}
|
||||
|
||||
invalidHosts := []string{
|
||||
"1:2:3",
|
||||
"[2001:db8::1",
|
||||
"2001:db8::1]",
|
||||
"[example.com]",
|
||||
"example..com",
|
||||
"-example.com",
|
||||
"example-.com",
|
||||
"exa_mple.com",
|
||||
"999.1.1.1",
|
||||
}
|
||||
for _, host := range invalidHosts {
|
||||
if _, _, err := normalizeTunnelProbeTarget(host, 443); err == nil {
|
||||
t.Fatalf("expected invalid host %q to be rejected", host)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseTunnelProbeTargetFromRequest(t *testing.T) {
|
||||
req := map[string]interface{}{
|
||||
"probeTargetHost": "speed.example.com",
|
||||
"probeTargetPort": float64(1443),
|
||||
}
|
||||
target, configured, err := parseTunnelProbeTargetFromRequest(req)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !configured || target.Host != "speed.example.com" || target.Port != 1443 {
|
||||
t.Fatalf("unexpected request target: %+v configured=%v", target, configured)
|
||||
}
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
@@ -15,7 +16,6 @@ const (
|
||||
tunnelQualityProbeInterval = 1 * time.Second
|
||||
tunnelQualityProbeTimeout = 8 * time.Second
|
||||
tunnelQualityPingTimeoutMs = 5000
|
||||
tunnelQualityRetention = 24 * time.Hour // keep 24h of history
|
||||
tunnelQualityPruneInterval = 10 * time.Minute
|
||||
tunnelQualityReportInterval = 30 * time.Second // DB save interval
|
||||
)
|
||||
@@ -42,6 +42,8 @@ type tunnelQualitySnapshot struct {
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
Timestamp int64 `json:"timestamp"`
|
||||
ChainDetails string `json:"chainDetails,omitempty"`
|
||||
ProbeTargetHost string `json:"probeTargetHost,omitempty"`
|
||||
ProbeTargetPort int `json:"probeTargetPort,omitempty"`
|
||||
|
||||
// internal fields for db reporting
|
||||
lastDBWrite int64 `json:"-"`
|
||||
@@ -57,6 +59,7 @@ type tunnelQualityProber struct {
|
||||
interval time.Duration
|
||||
lastPrune int64
|
||||
probing int32 // atomic flag: 1 = probeAll running, 0 = idle
|
||||
probeNode bestExitProbeFunc
|
||||
}
|
||||
|
||||
// newTunnelQualityProber creates a new prober (not yet running).
|
||||
@@ -128,12 +131,19 @@ func (p *tunnelQualityProber) isEnabled() bool {
|
||||
return p.handler.isTunnelQualityMonitoringEnabled()
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) retentionDays() int {
|
||||
if p == nil || p.handler == nil || p.handler.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
// maybePrune deletes old quality rows periodically (mirrors PruneServiceMonitorResults).
|
||||
func (p *tunnelQualityProber) maybePrune() {
|
||||
if !p.isEnabled() {
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if p.lastPrune > 0 && now-p.lastPrune < int64(tunnelQualityPruneInterval/time.Millisecond) {
|
||||
return
|
||||
@@ -145,7 +155,7 @@ func (p *tunnelQualityProber) maybePrune() {
|
||||
return
|
||||
}
|
||||
|
||||
cutoff := now - int64(tunnelQualityRetention/time.Millisecond)
|
||||
cutoff := now - int64(time.Duration(p.retentionDays())*24*time.Hour/time.Millisecond)
|
||||
if err := h.repo.PruneTunnelQualityResults(cutoff); err != nil {
|
||||
log.Printf("tunnel_quality_prober: prune err=%v", err)
|
||||
}
|
||||
@@ -219,6 +229,9 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
p.storeResult(snap)
|
||||
return
|
||||
}
|
||||
probeTarget := effectiveTunnelProbeTargetValues(tunnel.ProbeTargetHost, tunnel.ProbeTargetPort)
|
||||
snap.ProbeTargetHost = probeTarget.Host
|
||||
snap.ProbeTargetPort = probeTarget.Port
|
||||
|
||||
chainRows, err := h.listChainNodesForTunnel(tunnelID)
|
||||
if err != nil || len(chainRows) == 0 {
|
||||
@@ -235,12 +248,13 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
pingTimeoutMS: tunnelQualityPingTimeoutMs,
|
||||
timeoutMessage: "探测超时",
|
||||
}
|
||||
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget)
|
||||
|
||||
switch tunnel.Type {
|
||||
case 1:
|
||||
// Port forwarding: entry → Bing only
|
||||
// Port forwarding: entry → public probe target only.
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(inNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -287,7 +301,7 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
hops = append(hops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
|
||||
fromNode, _ := h.getNodeRecord(source.NodeID)
|
||||
targetIP, targetPort, resolveErr := resolveChainProbeTarget(fromNode, targetNode, target.Port, ipPreference, target.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
@@ -302,7 +316,7 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
hop.TargetIP = targetIP
|
||||
hop.TargetPort = targetPort
|
||||
|
||||
lat, loss, err := p.tcpPingNode(source.NodeID, targetIP, targetPort, options)
|
||||
lat, loss, err := p.pingNode(source.NodeID, targetIP, targetPort, options)
|
||||
if err == nil {
|
||||
hop.Latency = lat
|
||||
hop.Loss = loss
|
||||
@@ -338,7 +352,7 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
|
||||
// Exit → Bing
|
||||
if len(outNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(outNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
lat, loss, err := p.pingNode(outNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -352,9 +366,9 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
|
||||
snap.Success = probeOK
|
||||
default:
|
||||
// Unknown type: entry → Bing
|
||||
// Unknown type: entry → public probe target.
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.tcpPingNode(inNodes[0].NodeID, "www.bing.com", 443, options)
|
||||
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -368,6 +382,58 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
p.storeResult(snap)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget) {
|
||||
if p == nil || p.handler == nil || p.handler.bestExit == nil || len(outNodes) <= 1 {
|
||||
return
|
||||
}
|
||||
if !isBestTunnelStrategy(outNodes[0].Strategy) {
|
||||
return
|
||||
}
|
||||
owners := bestExitChainOwners(inNodes, chainHops)
|
||||
if len(owners) == 0 {
|
||||
return
|
||||
}
|
||||
nodeMap := make(map[int64]*nodeRecord, len(owners)+len(outNodes))
|
||||
for _, owner := range owners {
|
||||
if node, err := p.handler.getNodeRecord(owner.NodeID); err == nil && node != nil {
|
||||
nodeMap[owner.NodeID] = node
|
||||
}
|
||||
}
|
||||
for _, exit := range outNodes {
|
||||
if node, err := p.handler.getNodeRecord(exit.NodeID); err == nil && node != nil {
|
||||
nodeMap[exit.NodeID] = node
|
||||
}
|
||||
}
|
||||
// This best-exit decision cache is per decision round; the display-oriented
|
||||
// tunnel quality snapshot may still collect its own first-exit public probe.
|
||||
roundPinger := newBestExitRoundPinger(p.pingNode)
|
||||
for _, owner := range owners {
|
||||
if nodeMap[owner.NodeID] == nil {
|
||||
continue
|
||||
}
|
||||
key := bestExitOwnerKey{TunnelID: tunnelID, OwnerNodeID: owner.NodeID}
|
||||
p.handler.bestExit.ensureApplied(key, outNodes[0].NodeID, time.Now())
|
||||
scores := evaluateBestExitOwner(owner, outNodes, nodeMap, ipPreference, options, probeTarget, roundPinger)
|
||||
decision := p.handler.bestExit.observeScores(key, scores, time.Now())
|
||||
if decision.Switch {
|
||||
now := time.Now()
|
||||
if err := p.handler.applyBestExitChainOrder(tunnelID, owner.NodeID, outNodes, decision.Scores, ipPreference); err != nil {
|
||||
log.Printf("best_exit: switch apply failed tunnel=%d owner=%d exit=%d err=%v", tunnelID, owner.NodeID, decision.ExitNodeID, err)
|
||||
p.handler.bestExit.recordApplyFailure(key, decision.ExitNodeID, now)
|
||||
continue
|
||||
}
|
||||
p.handler.bestExit.setApplied(key, decision.ExitNodeID, time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) pingNode(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
if p != nil && p.probeNode != nil {
|
||||
return p.probeNode(nodeID, ip, port, options)
|
||||
}
|
||||
return p.tcpPingNode(nodeID, ip, port, options)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) tcpPingNode(nodeID int64, ip string, port int, options diagnosisExecOptions) (latency float64, loss float64, err error) {
|
||||
h := p.handler
|
||||
if h == nil {
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 77, "quality-target", "speed.example.com", 8443)
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(30, 'exit-a', 'exit-secret', '10.0.0.30', '10.0.0.30', '', '30000-30010', '', 'v1', 1, 1, 1, ?, ?, 1, '[::]', '[::]', 0)
|
||||
`, time.Now().UnixMilli(), time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("insert exit node: %v", err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(77, '3', 30, 30001, 'round', 1, 'tls')
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert exit chain: %v", err)
|
||||
}
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
var calls []string
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
p.probeTunnel(77)
|
||||
|
||||
if !slices.Contains(calls, "10|speed.example.com|8443") {
|
||||
t.Fatalf("expected type 1 public probe from entry to configured target, calls=%+v", calls)
|
||||
}
|
||||
if slices.Contains(calls, "30|speed.example.com|8443") {
|
||||
t.Fatalf("did not expect type 1 public probe from exit node, calls=%+v", calls)
|
||||
}
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].ProbeTargetHost != "speed.example.com" || snaps[0].ProbeTargetPort != 8443 {
|
||||
t.Fatalf("unexpected snapshot target metadata: %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 78, "quality-target-incomplete", "speed.example.com", 8443)
|
||||
if err := h.repo.DB().Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, 78).Error; err != nil {
|
||||
t.Fatalf("delete chain rows: %v", err)
|
||||
}
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
p.probeTunnel(78)
|
||||
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].ErrorMessage == "" {
|
||||
t.Fatalf("expected incomplete chain error, got %+v", snaps[0])
|
||||
}
|
||||
if snaps[0].ProbeTargetHost != "speed.example.com" || snaps[0].ProbeTargetPort != 8443 {
|
||||
t.Fatalf("unexpected snapshot target metadata: %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
@@ -13,6 +13,13 @@ import (
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
// failedForward tracks a forward that failed redeployment, for retry.
|
||||
type failedForward struct {
|
||||
id int64
|
||||
forward *forwardRecord
|
||||
err error
|
||||
}
|
||||
|
||||
const (
|
||||
githubRepo = "Sagit-chu/flvx"
|
||||
githubAPIBase = "https://api.github.com"
|
||||
@@ -32,6 +39,8 @@ var (
|
||||
testKeywordPattern = regexp.MustCompile(`(?i)(alpha|beta|rc)`)
|
||||
)
|
||||
|
||||
const nodeOnlineRedeployCooldown = 30 * time.Second
|
||||
|
||||
type githubRelease struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
@@ -389,22 +398,123 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
|
||||
}
|
||||
|
||||
func (h *Handler) onNodeOnline(nodeID int64) {
|
||||
h.consumeNodePendingUpgradeRedeploy(nodeID)
|
||||
h.redeployNodeRuntimeAfterUpgrade(nodeID)
|
||||
if !h.startNodeOnlineRedeploy(nodeID, time.Now()) {
|
||||
return
|
||||
}
|
||||
defer h.finishNodeOnlineRedeploy(nodeID)
|
||||
|
||||
// Reconcile node runtime on the first reconnect, but suppress rapid flapping
|
||||
// so websocket churn does not trigger repeated full redeploy storms.
|
||||
if !h.redeployNodeRuntimeAfterUpgrade(nodeID) {
|
||||
h.markNodePendingUpgradeRedeploy(nodeID)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
func (h *Handler) startNodeOnlineRedeploy(nodeID int64, now time.Time) bool {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return false
|
||||
}
|
||||
if now.IsZero() {
|
||||
now = time.Now()
|
||||
}
|
||||
|
||||
h.upgradeMu.Lock()
|
||||
defer h.upgradeMu.Unlock()
|
||||
if h.pendingUpgradeRedeploy == nil {
|
||||
h.pendingUpgradeRedeploy = make(map[int64]struct{})
|
||||
}
|
||||
if h.nodeOnlineRedeployAt == nil {
|
||||
h.nodeOnlineRedeployAt = make(map[int64]time.Time)
|
||||
}
|
||||
if h.nodeOnlineRedeployQueued == nil {
|
||||
h.nodeOnlineRedeployQueued = make(map[int64]struct{})
|
||||
}
|
||||
if h.nodeOnlineRedeploying == nil {
|
||||
h.nodeOnlineRedeploying = make(map[int64]struct{})
|
||||
}
|
||||
|
||||
_, pendingUpgrade := h.pendingUpgradeRedeploy[nodeID]
|
||||
lastRedeployAt := h.nodeOnlineRedeployAt[nodeID]
|
||||
_, inFlight := h.nodeOnlineRedeploying[nodeID]
|
||||
if fireAt, start := nextNodeOnlineRedeployFireAt(lastRedeployAt, now, pendingUpgrade, inFlight); !start {
|
||||
h.queueNodeOnlineRedeployLocked(nodeID, fireAt)
|
||||
return false
|
||||
}
|
||||
|
||||
delete(h.pendingUpgradeRedeploy, nodeID)
|
||||
h.nodeOnlineRedeployAt[nodeID] = now
|
||||
h.nodeOnlineRedeploying[nodeID] = struct{}{}
|
||||
return true
|
||||
}
|
||||
|
||||
func nextNodeOnlineRedeployFireAt(lastRedeployAt, now time.Time, pendingUpgrade bool, inFlight bool) (time.Time, bool) {
|
||||
if now.IsZero() {
|
||||
now = time.Now()
|
||||
}
|
||||
if inFlight {
|
||||
fireAt := now.Add(nodeOnlineRedeployCooldown)
|
||||
if !lastRedeployAt.IsZero() {
|
||||
cooldownAt := lastRedeployAt.Add(nodeOnlineRedeployCooldown)
|
||||
if cooldownAt.After(now) {
|
||||
fireAt = cooldownAt
|
||||
}
|
||||
}
|
||||
return fireAt, false
|
||||
}
|
||||
if !pendingUpgrade && !lastRedeployAt.IsZero() && now.Sub(lastRedeployAt) < nodeOnlineRedeployCooldown {
|
||||
return lastRedeployAt.Add(nodeOnlineRedeployCooldown), false
|
||||
}
|
||||
return time.Time{}, true
|
||||
}
|
||||
|
||||
func (h *Handler) queueNodeOnlineRedeployLocked(nodeID int64, fireAt time.Time) {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return
|
||||
}
|
||||
if h.nodeOnlineRedeployQueued == nil {
|
||||
h.nodeOnlineRedeployQueued = make(map[int64]struct{})
|
||||
}
|
||||
if _, queued := h.nodeOnlineRedeployQueued[nodeID]; queued {
|
||||
return
|
||||
}
|
||||
if fireAt.IsZero() {
|
||||
fireAt = time.Now().Add(nodeOnlineRedeployCooldown)
|
||||
}
|
||||
delay := time.Until(fireAt)
|
||||
if delay < 0 {
|
||||
delay = 0
|
||||
}
|
||||
h.nodeOnlineRedeployQueued[nodeID] = struct{}{}
|
||||
time.AfterFunc(delay, func() {
|
||||
h.upgradeMu.Lock()
|
||||
delete(h.nodeOnlineRedeployQueued, nodeID)
|
||||
h.upgradeMu.Unlock()
|
||||
h.onNodeOnline(nodeID)
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) finishNodeOnlineRedeploy(nodeID int64) {
|
||||
if h == nil || nodeID <= 0 {
|
||||
return
|
||||
}
|
||||
h.upgradeMu.Lock()
|
||||
delete(h.nodeOnlineRedeploying, nodeID)
|
||||
h.upgradeMu.Unlock()
|
||||
}
|
||||
|
||||
func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) bool {
|
||||
tunnelIDs, err := h.repo.ListActiveTunnelIDsByNode(nodeID)
|
||||
if err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: list tunnels for node %d failed: %v\n", nodeID, err)
|
||||
return
|
||||
return false
|
||||
}
|
||||
forwardIDs, err := h.repo.ListForwardIDsByNode(nodeID)
|
||||
if err != nil {
|
||||
fmt.Printf("post-upgrade redeploy: list forwards for node %d failed: %v\n", nodeID, err)
|
||||
return
|
||||
return false
|
||||
}
|
||||
|
||||
// First pass: deploy everything
|
||||
tunnelFailed := make(map[int64]struct{})
|
||||
for _, tunnelID := range tunnelIDs {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
@@ -413,6 +523,9 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
// Collect forwards that failed independently (not skipped due to tunnel failure)
|
||||
var failedForwards []failedForward
|
||||
|
||||
for _, forwardID := range forwardIDs {
|
||||
forward, getErr := h.getForwardRecord(forwardID)
|
||||
if getErr != nil || forward == nil {
|
||||
@@ -422,7 +535,87 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
failedForwards = append(failedForwards, failedForward{id: forwardID, forward: forward, err: err})
|
||||
fmt.Printf("post-upgrade redeploy: forward %d failed on node %d: %v\n", forwardID, nodeID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Retry failed items with exponential backoff (max 3 attempts)
|
||||
return h.retryFailedRedeploys(nodeID, tunnelFailed, failedForwards)
|
||||
}
|
||||
|
||||
// isRetryableError returns true if the error looks transient and worth retrying.
|
||||
func isRetryableError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
// Skip non-retryable errors: not-found, already-exists, validation errors
|
||||
if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(msg, "already exists") || strings.Contains(msg, "已存在") {
|
||||
return false
|
||||
}
|
||||
// Everything else (timeout, connection lost, port in use, etc.) is retryable
|
||||
return true
|
||||
}
|
||||
|
||||
// retryFailedRedeploys retries failed tunnels and forwards with exponential backoff.
|
||||
func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]struct{}, failedForwards []failedForward) bool {
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
const maxRetries = 3
|
||||
baseDelay := time.Second
|
||||
|
||||
for attempt := 1; attempt <= maxRetries; attempt++ {
|
||||
delay := baseDelay * time.Duration(1<<uint(attempt-1)) // 1s, 2s, 4s
|
||||
time.Sleep(delay)
|
||||
|
||||
// Retry failed tunnels
|
||||
for tunnelID := range tunnelFailed {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err == nil {
|
||||
delete(tunnelFailed, tunnelID)
|
||||
fmt.Printf("post-upgrade redeploy retry: tunnel %d succeeded on node %d (attempt %d)\n", tunnelID, nodeID, attempt)
|
||||
} else if !isRetryableError(err) {
|
||||
delete(tunnelFailed, tunnelID) // Non-retryable, don't retry again
|
||||
} else {
|
||||
fmt.Printf("post-upgrade redeploy retry: tunnel %d still failing on node %d (attempt %d): %v\n", tunnelID, nodeID, attempt, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Retry failed forwards
|
||||
var stillFailed []failedForward
|
||||
for _, ff := range failedForwards {
|
||||
if _, skipped := tunnelFailed[ff.forward.TunnelID]; skipped {
|
||||
stillFailed = append(stillFailed, ff) // Tunnel still failed, skip forward
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(ff.forward, "UpdateService", true); err == nil {
|
||||
fmt.Printf("post-upgrade redeploy retry: forward %d succeeded on node %d (attempt %d)\n", ff.id, nodeID, attempt)
|
||||
} else if !isRetryableError(err) {
|
||||
// Non-retryable, drop it
|
||||
} else {
|
||||
stillFailed = append(stillFailed, ff)
|
||||
fmt.Printf("post-upgrade redeploy retry: forward %d still failing on node %d (attempt %d): %v\n", ff.id, nodeID, attempt, err)
|
||||
}
|
||||
}
|
||||
failedForwards = stillFailed
|
||||
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
fmt.Printf("post-upgrade redeploy retry: all items recovered on node %d\n", nodeID)
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// Final summary
|
||||
for tunnelID := range tunnelFailed {
|
||||
fmt.Printf("post-upgrade redeploy: tunnel %d permanently failed on node %d after retries\n", tunnelID, nodeID)
|
||||
}
|
||||
for _, ff := range failedForwards {
|
||||
fmt.Printf("post-upgrade redeploy: forward %d permanently failed on node %d after retries\n", ff.id, nodeID)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestStartNodeOnlineRedeploySkipsRecentReconnects(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
|
||||
if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
|
||||
t.Fatalf("expected recent reconnect to skip redeploy")
|
||||
}
|
||||
if h.consumeNodePendingUpgradeRedeploy(54) {
|
||||
t.Fatalf("did not expect pending upgrade marker to be consumed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartNodeOnlineRedeployAllowsPendingUpgradeDuringCooldown(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
h.markNodePendingUpgradeRedeploy(54)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
|
||||
t.Fatalf("expected pending upgrade reconnect to bypass cooldown")
|
||||
}
|
||||
if h.consumeNodePendingUpgradeRedeploy(54) {
|
||||
t.Fatalf("expected pending upgrade marker to be consumed during redeploy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartNodeOnlineRedeployQueuesCooldownReconnect(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
|
||||
if h.startNodeOnlineRedeploy(54, now.Add(5*time.Second)) {
|
||||
t.Fatalf("expected cooldown reconnect to skip immediate redeploy")
|
||||
}
|
||||
if _, queued := h.nodeOnlineRedeployQueued[54]; !queued {
|
||||
t.Fatalf("expected cooldown reconnect to queue a follow-up redeploy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartNodeOnlineRedeployKeepsPendingUpgradeWhileInFlight(t *testing.T) {
|
||||
h := &Handler{
|
||||
pendingUpgradeRedeploy: map[int64]struct{}{},
|
||||
nodeOnlineRedeployAt: map[int64]time.Time{},
|
||||
nodeOnlineRedeployQueued: map[int64]struct{}{},
|
||||
nodeOnlineRedeploying: map[int64]struct{}{},
|
||||
}
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
|
||||
if !h.startNodeOnlineRedeploy(54, now) {
|
||||
t.Fatalf("expected first reconnect to redeploy")
|
||||
}
|
||||
h.markNodePendingUpgradeRedeploy(54)
|
||||
|
||||
if h.startNodeOnlineRedeploy(54, now.Add(time.Second)) {
|
||||
t.Fatalf("expected in-flight redeploy to suppress parallel restart")
|
||||
}
|
||||
if !h.consumeNodePendingUpgradeRedeploy(54) {
|
||||
t.Fatalf("expected pending upgrade marker to remain for the next retry")
|
||||
}
|
||||
h.finishNodeOnlineRedeploy(54)
|
||||
}
|
||||
|
||||
func TestNextNodeOnlineRedeployFireAtDefersExpiredInFlightReconnect(t *testing.T) {
|
||||
now := time.Unix(1_777_176_720, 0)
|
||||
last := now.Add(-nodeOnlineRedeployCooldown - 5*time.Second)
|
||||
|
||||
fireAt, start := nextNodeOnlineRedeployFireAt(last, now, false, true)
|
||||
if start {
|
||||
t.Fatalf("expected in-flight reconnect to queue instead of starting immediately")
|
||||
}
|
||||
|
||||
want := now.Add(nodeOnlineRedeployCooldown)
|
||||
if !fireAt.Equal(want) {
|
||||
t.Fatalf("expected queued reconnect at %s, got %s", want, fireAt)
|
||||
}
|
||||
}
|
||||
@@ -105,6 +105,10 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/system/") {
|
||||
return true
|
||||
}
|
||||
|
||||
if strings.HasPrefix(path, "/api/v1/group/") {
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/monitoring"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
@@ -31,7 +32,6 @@ type IngestionService struct {
|
||||
nodeBuffer []*model.NodeMetric
|
||||
nodeBufferMu sync.Mutex
|
||||
flushInterval time.Duration
|
||||
retentionDays int
|
||||
}
|
||||
|
||||
func NewIngestionService(repo *repo.Repository) *IngestionService {
|
||||
@@ -39,7 +39,6 @@ func NewIngestionService(repo *repo.Repository) *IngestionService {
|
||||
repo: repo,
|
||||
nodeBuffer: make([]*model.NodeMetric, 0, 500),
|
||||
flushInterval: 30 * time.Second,
|
||||
retentionDays: 7,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -111,7 +110,22 @@ func (s *IngestionService) flushNodeMetrics() {
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetrics() {
|
||||
cutoff := time.Now().Add(-time.Duration(s.retentionDays) * 24 * time.Hour).UnixMilli()
|
||||
s.pruneMetricsAt(time.Now())
|
||||
}
|
||||
|
||||
func (s *IngestionService) retentionDaysFromConfig() int {
|
||||
if s == nil || s.repo == nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
cfg, err := s.repo.GetConfigsByNames([]string{monitoring.ConfigMonitorRetentionDays})
|
||||
if err != nil {
|
||||
return monitoring.DefaultMonitorRetentionDays
|
||||
}
|
||||
return monitoring.MonitoringRetentionDaysFromConfigMap(cfg)
|
||||
}
|
||||
|
||||
func (s *IngestionService) pruneMetricsAt(now time.Time) {
|
||||
cutoff := now.Add(-time.Duration(s.retentionDaysFromConfig()) * 24 * time.Hour).UnixMilli()
|
||||
if s.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
@@ -215,7 +216,6 @@ func TestPruneMetrics(t *testing.T) {
|
||||
defer r.Close()
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.retentionDays = 1
|
||||
|
||||
info := SystemInfo{CPUUsage: 50.0, MemoryUsage: 60.0, DiskUsage: 30.0}
|
||||
|
||||
@@ -233,6 +233,39 @@ func TestPruneMetrics(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPruneMetricsUsesConfiguredRetentionDays(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.UpsertConfig("monitor_retention_days", "2", now); err != nil {
|
||||
t.Fatalf("upsert retention config: %v", err)
|
||||
}
|
||||
|
||||
oldMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(3*24*time.Hour/time.Millisecond), CPUUsage: 10}
|
||||
newMetric := &model.NodeMetric{NodeID: 1, Timestamp: now - int64(1*24*time.Hour/time.Millisecond), CPUUsage: 20}
|
||||
if err := r.InsertNodeMetric(oldMetric); err != nil {
|
||||
t.Fatalf("insert old metric: %v", err)
|
||||
}
|
||||
if err := r.InsertNodeMetric(newMetric); err != nil {
|
||||
t.Fatalf("insert new metric: %v", err)
|
||||
}
|
||||
|
||||
svc := NewIngestionService(r)
|
||||
svc.pruneMetricsAt(time.UnixMilli(now))
|
||||
|
||||
metrics, err := r.GetNodeMetrics(1, now-int64(4*24*time.Hour/time.Millisecond), now+1000)
|
||||
if err != nil {
|
||||
t.Fatalf("get node metrics: %v", err)
|
||||
}
|
||||
if len(metrics) != 1 || metrics[0].CPUUsage != 20 {
|
||||
t.Fatalf("expected only newer metric to remain, got %#v", metrics)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultipleNodes(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
ConfigMonitorRetentionDays = "monitor_retention_days"
|
||||
DefaultMonitorRetentionDays = 7
|
||||
MinMonitorRetentionDays = 1
|
||||
MaxMonitorRetentionDays = 3650
|
||||
)
|
||||
|
||||
func MonitoringRetentionDaysFromConfigMap(cfg map[string]string) int {
|
||||
if cfg == nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
days, err := parseMonitoringRetentionDays(cfg[ConfigMonitorRetentionDays])
|
||||
if err != nil {
|
||||
return DefaultMonitorRetentionDays
|
||||
}
|
||||
return days
|
||||
}
|
||||
|
||||
func NormalizeMonitoringRetentionDays(value string) (string, error) {
|
||||
days, err := parseMonitoringRetentionDays(value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strconv.Itoa(days), nil
|
||||
}
|
||||
|
||||
func parseMonitoringRetentionDays(value string) (int, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return 0, fmt.Errorf("监控数据保留天数不能为空")
|
||||
}
|
||||
days, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须是整数")
|
||||
}
|
||||
if days < MinMonitorRetentionDays || days > MaxMonitorRetentionDays {
|
||||
return 0, fmt.Errorf("监控数据保留天数必须在 %d 到 %d 之间", MinMonitorRetentionDays, MaxMonitorRetentionDays)
|
||||
}
|
||||
return days, nil
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package monitoring
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestMonitoringRetentionDaysFromConfigMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg map[string]string
|
||||
want int
|
||||
}{
|
||||
{"missing uses default", nil, 7},
|
||||
{"valid custom", map[string]string{ConfigMonitorRetentionDays: "3"}, 3},
|
||||
{"trimmed custom", map[string]string{ConfigMonitorRetentionDays: " 30 "}, 30},
|
||||
{"invalid uses default", map[string]string{ConfigMonitorRetentionDays: "abc"}, 7},
|
||||
{"too small uses default", map[string]string{ConfigMonitorRetentionDays: "0"}, 7},
|
||||
{"too large uses default", map[string]string{ConfigMonitorRetentionDays: "3651"}, 7},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := MonitoringRetentionDaysFromConfigMap(tc.cfg); got != tc.want {
|
||||
t.Fatalf("expected %d, got %d", tc.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeMonitoringRetentionDays(t *testing.T) {
|
||||
for _, value := range []string{"1", "7", "3650", " 30 "} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err != nil || got == "" {
|
||||
t.Fatalf("expected %q valid, got value=%q err=%v", value, got, err)
|
||||
}
|
||||
}
|
||||
|
||||
for _, value := range []string{"", "0", "-1", "3651", "abc", "1.5"} {
|
||||
if got, err := NormalizeMonitoringRetentionDays(value); err == nil {
|
||||
t.Fatalf("expected %q invalid, got value=%q", value, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -30,21 +30,24 @@ func (User) TableName() string { return "user" }
|
||||
|
||||
// Forward maps to the "forward" table.
|
||||
type Forward struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
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" }
|
||||
@@ -116,18 +119,20 @@ type StatisticsFlow struct {
|
||||
func (StatisticsFlow) TableName() string { return "statistics_flow" }
|
||||
|
||||
type Tunnel struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio;not null;default:1.0"`
|
||||
Type int `gorm:"not null"`
|
||||
Protocol string `gorm:"type:varchar(10);not null;default:'tls'"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
InIP sql.NullString `gorm:"column:in_ip;type:text"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
IPPreference string `gorm:"column:ip_preference;type:varchar(10);not null;default:''"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio;not null;default:1.0"`
|
||||
Type int `gorm:"not null"`
|
||||
Protocol string `gorm:"type:varchar(10);not null;default:'tls'"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
InIP sql.NullString `gorm:"column:in_ip;type:text"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
IPPreference string `gorm:"column:ip_preference;type:varchar(10);not null;default:''"`
|
||||
ProbeTargetHost string `gorm:"column:probe_target_host;type:text;not null;default:''"`
|
||||
ProbeTargetPort int `gorm:"column:probe_target_port;not null;default:0"`
|
||||
}
|
||||
|
||||
func (Tunnel) TableName() string { return "tunnel" }
|
||||
@@ -400,19 +405,21 @@ type NodeBackup struct {
|
||||
}
|
||||
|
||||
type TunnelBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
TrafficRatio float64 `json:"trafficRatio"`
|
||||
Type int `json:"type"`
|
||||
Protocol string `json:"protocol"`
|
||||
Flow int64 `json:"flow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
InIP string `json:"inIp,omitempty"`
|
||||
Inx int `json:"inx"`
|
||||
IPPreference string `json:"ipPreference,omitempty"`
|
||||
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
TrafficRatio float64 `json:"trafficRatio"`
|
||||
Type int `json:"type"`
|
||||
Protocol string `json:"protocol"`
|
||||
Flow int64 `json:"flow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
InIP string `json:"inIp,omitempty"`
|
||||
Inx int `json:"inx"`
|
||||
IPPreference string `json:"ipPreference,omitempty"`
|
||||
ProbeTargetHost string `json:"probeTargetHost,omitempty"`
|
||||
ProbeTargetPort int `json:"probeTargetPort,omitempty"`
|
||||
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
|
||||
}
|
||||
|
||||
type ChainTunnelBackup struct {
|
||||
@@ -427,21 +434,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 {
|
||||
@@ -530,26 +540,31 @@ type ImportResult struct {
|
||||
|
||||
// ForwardRecord is a minimal forward view used by control plane and flow policy.
|
||||
type ForwardRecord struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
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.
|
||||
type TunnelRecord struct {
|
||||
ID int64
|
||||
Type int
|
||||
Status int
|
||||
Flow int64
|
||||
TrafficRatio float64
|
||||
Protocol string
|
||||
ID int64
|
||||
Type int
|
||||
Status int
|
||||
Flow int64
|
||||
TrafficRatio float64
|
||||
Protocol string
|
||||
ProbeTargetHost string
|
||||
ProbeTargetPort int
|
||||
}
|
||||
|
||||
type UserQuotaView struct {
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -58,7 +65,16 @@ type TunnelQuality = model.TunnelQuality
|
||||
// ─── Repository ──────────────────────────────────────────────────────
|
||||
|
||||
type Repository struct {
|
||||
db *gorm.DB
|
||||
db *gorm.DB
|
||||
dbPath string
|
||||
}
|
||||
|
||||
type FlowUploadCounterDelta struct {
|
||||
ForwardID int64
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
}
|
||||
|
||||
func (r *Repository) DB() *gorm.DB {
|
||||
@@ -68,6 +84,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) {
|
||||
@@ -110,7 +199,7 @@ func Open(path string) (*Repository, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Repository{db: db}, nil
|
||||
return &Repository{db: db, dbPath: path}, nil
|
||||
}
|
||||
|
||||
func OpenPostgres(dsn string) (*Repository, error) {
|
||||
@@ -129,6 +218,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 +244,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
|
||||
@@ -204,6 +304,7 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
m := db.Migrator()
|
||||
hasNode := m.HasTable(&model.Node{})
|
||||
hasTunnel := m.HasTable(&model.Tunnel{})
|
||||
hasForward := m.HasTable(&model.Forward{})
|
||||
|
||||
for _, item := range models {
|
||||
if hasNode {
|
||||
@@ -216,6 +317,11 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if hasForward {
|
||||
if _, ok := item.(*model.Forward); ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if err := db.AutoMigrate(item); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -285,7 +391,7 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
}
|
||||
|
||||
if m.HasTable(&model.Tunnel{}) {
|
||||
for _, field := range []string{"Inx", "IPPreference"} {
|
||||
for _, field := range []string{"Inx", "IPPreference", "ProbeTargetHost", "ProbeTargetPort"} {
|
||||
if m.HasColumn(&model.Tunnel{}, field) {
|
||||
continue
|
||||
}
|
||||
@@ -295,6 +401,17 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
}
|
||||
}
|
||||
|
||||
if m.HasTable(&model.Forward{}) {
|
||||
for _, field := range []string{"MaxConn", "IPMaxConn", "IPSpeedID", "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 +796,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 +831,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 +872,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 +919,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
|
||||
@@ -1014,11 +1147,13 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
|
||||
"id": t.ID, "inx": t.Inx, "name": t.Name,
|
||||
"type": t.Type, "flow": t.Flow, "trafficRatio": t.TrafficRatio,
|
||||
"status": t.Status, "createdTime": t.CreatedTime,
|
||||
"inIp": nullableString(t.InIP),
|
||||
"ipPreference": t.IPPreference,
|
||||
"inNodeId": make([]map[string]interface{}, 0),
|
||||
"outNodeId": make([]map[string]interface{}, 0),
|
||||
"chainNodes": make([][]map[string]interface{}, 0),
|
||||
"inIp": nullableString(t.InIP),
|
||||
"ipPreference": t.IPPreference,
|
||||
"probeTargetHost": t.ProbeTargetHost,
|
||||
"probeTargetPort": t.ProbeTargetPort,
|
||||
"inNodeId": make([]map[string]interface{}, 0),
|
||||
"outNodeId": make([]map[string]interface{}, 0),
|
||||
"chainNodes": make([][]map[string]interface{}, 0),
|
||||
}
|
||||
orderedIDs = append(orderedIDs, t.ID)
|
||||
}
|
||||
@@ -1934,6 +2069,7 @@ func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) {
|
||||
Type: t.Type, Protocol: t.Protocol, Flow: t.Flow,
|
||||
CreatedTime: t.CreatedTime, UpdatedTime: t.UpdatedTime,
|
||||
Status: t.Status, Inx: t.Inx, IPPreference: t.IPPreference,
|
||||
ProbeTargetHost: t.ProbeTargetHost, ProbeTargetPort: t.ProbeTargetPort,
|
||||
}
|
||||
if t.InIP.Valid {
|
||||
b.InIP = t.InIP.String
|
||||
@@ -1987,6 +2123,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 {
|
||||
@@ -2319,23 +2465,25 @@ func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, e
|
||||
count := 0
|
||||
for _, t := range tunnels {
|
||||
item := model.Tunnel{
|
||||
ID: t.ID,
|
||||
Name: t.Name,
|
||||
TrafficRatio: t.TrafficRatio,
|
||||
Type: t.Type,
|
||||
Protocol: t.Protocol,
|
||||
Flow: t.Flow,
|
||||
CreatedTime: t.CreatedTime,
|
||||
UpdatedTime: now,
|
||||
Status: t.Status,
|
||||
InIP: sql.NullString{String: t.InIP, Valid: true},
|
||||
Inx: t.Inx,
|
||||
IPPreference: t.IPPreference,
|
||||
ID: t.ID,
|
||||
Name: t.Name,
|
||||
TrafficRatio: t.TrafficRatio,
|
||||
Type: t.Type,
|
||||
Protocol: t.Protocol,
|
||||
Flow: t.Flow,
|
||||
CreatedTime: t.CreatedTime,
|
||||
UpdatedTime: now,
|
||||
Status: t.Status,
|
||||
InIP: sql.NullString{String: t.InIP, Valid: true},
|
||||
Inx: t.Inx,
|
||||
IPPreference: t.IPPreference,
|
||||
ProbeTargetHost: t.ProbeTargetHost,
|
||||
ProbeTargetPort: t.ProbeTargetPort,
|
||||
}
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"name", "traffic_ratio", "type", "protocol", "flow", "updated_time", "status", "in_ip", "inx", "ip_preference",
|
||||
"name", "traffic_ratio", "type", "protocol", "flow", "updated_time", "status", "in_ip", "inx", "ip_preference", "probe_target_host", "probe_target_port",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
if err != nil {
|
||||
@@ -2367,29 +2515,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 +3496,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 {
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestBackupRoundTripsTunnelProbeTarget(t *testing.T) {
|
||||
source, err := Open(filepath.Join(t.TempDir(), "source.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open source repo: %v", err)
|
||||
}
|
||||
defer source.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := source.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, probe_target_host, probe_target_port)
|
||||
VALUES(20, 'backup-target', 1, 2, 'tls', 1, ?, ?, 1, '', 1, 'speed.example.com', 8443)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert source tunnel: %v", err)
|
||||
}
|
||||
|
||||
backup, err := source.ExportAll()
|
||||
if err != nil {
|
||||
t.Fatalf("export backup: %v", err)
|
||||
}
|
||||
if len(backup.Tunnels) != 1 {
|
||||
t.Fatalf("expected one exported tunnel, got %d", len(backup.Tunnels))
|
||||
}
|
||||
if backup.Tunnels[0].ProbeTargetHost != "speed.example.com" || backup.Tunnels[0].ProbeTargetPort != 8443 {
|
||||
t.Fatalf("unexpected exported probe target: %+v", backup.Tunnels[0])
|
||||
}
|
||||
|
||||
dest, err := Open(filepath.Join(t.TempDir(), "dest.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open dest repo: %v", err)
|
||||
}
|
||||
defer dest.Close()
|
||||
|
||||
result, err := dest.Import(backup, []string{"tunnels"})
|
||||
if err != nil {
|
||||
t.Fatalf("import backup: %v", err)
|
||||
}
|
||||
if result.TunnelsImported != 1 {
|
||||
t.Fatalf("expected one imported tunnel, got %d", result.TunnelsImported)
|
||||
}
|
||||
|
||||
items, err := dest.ListTunnels()
|
||||
if err != nil {
|
||||
t.Fatalf("list imported tunnels: %v", err)
|
||||
}
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("expected one imported tunnel item, got %d", len(items))
|
||||
}
|
||||
if items[0]["probeTargetHost"] != "speed.example.com" || items[0]["probeTargetPort"] != 8443 {
|
||||
t.Fatalf("unexpected imported probe target: %+v", items[0])
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
@@ -142,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")
|
||||
@@ -169,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")
|
||||
|
||||
@@ -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,16 +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,
|
||||
MaxConn: f.MaxConn,
|
||||
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"
|
||||
@@ -154,12 +253,14 @@ func (r *Repository) GetTunnelRecord(tunnelID int64) (*model.TunnelRecord, error
|
||||
return nil, err
|
||||
}
|
||||
tr := model.TunnelRecord{
|
||||
ID: t.ID,
|
||||
Type: t.Type,
|
||||
Status: t.Status,
|
||||
Flow: t.Flow,
|
||||
TrafficRatio: t.TrafficRatio,
|
||||
Protocol: t.Protocol,
|
||||
ID: t.ID,
|
||||
Type: t.Type,
|
||||
Status: t.Status,
|
||||
Flow: t.Flow,
|
||||
TrafficRatio: t.TrafficRatio,
|
||||
Protocol: t.Protocol,
|
||||
ProbeTargetHost: t.ProbeTargetHost,
|
||||
ProbeTargetPort: t.ProbeTargetPort,
|
||||
}
|
||||
if tr.Flow <= 0 {
|
||||
tr.Flow = 1
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
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 TestGetTunnelRecordIncludesProbeTarget(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "tunnel-record-probe-target.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx, probe_target_host, probe_target_port)
|
||||
VALUES(1, 't1', 1, 2, 'tls', 1, ?, ?, 1, NULL, 0, 'speed.example.com', 8443)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
|
||||
record, err := r.GetTunnelRecord(1)
|
||||
if err != nil {
|
||||
t.Fatalf("get tunnel record: %v", err)
|
||||
}
|
||||
if record == nil {
|
||||
t.Fatalf("expected tunnel record")
|
||||
}
|
||||
if record.ProbeTargetHost != "speed.example.com" || record.ProbeTargetPort != 8443 {
|
||||
t.Fatalf("unexpected probe target on record: %#v", record)
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package repo
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
@@ -58,6 +59,126 @@ func TestPrepareSQLiteLegacyColumnsAddsNodeMetadataColumns(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenBackfillsSQLiteLegacyTunnelProbeTargetColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "legacy.db")
|
||||
db, err := gorm.Open(gsqlite.Open(dbPath), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open legacy sqlite: %v", err)
|
||||
}
|
||||
|
||||
if err := db.Exec(`
|
||||
CREATE TABLE tunnel (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
traffic_ratio REAL NOT NULL DEFAULT 1.0,
|
||||
type INTEGER NOT NULL,
|
||||
protocol VARCHAR(10) NOT NULL DEFAULT 'tls',
|
||||
flow INTEGER NOT NULL,
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER NOT NULL,
|
||||
status INTEGER NOT NULL,
|
||||
in_ip TEXT
|
||||
)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create legacy tunnel table: %v", err)
|
||||
}
|
||||
if err := db.Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip)
|
||||
VALUES(1, 'legacy-tunnel', 1, 1, 'tls', 1, 1, 1, 1, '')
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert legacy tunnel: %v", err)
|
||||
}
|
||||
if sqlDB, _ := db.DB(); sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open migrated sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
m := r.DB().Migrator()
|
||||
for _, field := range []string{"ProbeTargetHost", "ProbeTargetPort"} {
|
||||
if !m.HasColumn(&model.Tunnel{}, field) {
|
||||
t.Fatalf("expected tunnel.%s column to exist", field)
|
||||
}
|
||||
}
|
||||
|
||||
var host string
|
||||
var port int
|
||||
if err := r.DB().Raw(`SELECT probe_target_host, probe_target_port FROM tunnel WHERE id = 1`).Row().Scan(&host, &port); err != nil {
|
||||
t.Fatalf("query probe target defaults: %v", err)
|
||||
}
|
||||
if host != "" || port != 0 {
|
||||
t.Fatalf("expected default probe target empty/0, got %q/%d", host, port)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenBackfillsSQLiteLegacyForwardColumns(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "legacy-forward.db")
|
||||
db, err := gorm.Open(gsqlite.Open(dbPath), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open legacy sqlite: %v", err)
|
||||
}
|
||||
|
||||
if err := db.Exec(`
|
||||
CREATE TABLE forward (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
user_id INTEGER NOT NULL,
|
||||
user_name VARCHAR(100) NOT NULL,
|
||||
name VARCHAR(100) NOT NULL,
|
||||
tunnel_id INTEGER NOT NULL,
|
||||
remote_addr TEXT NOT NULL,
|
||||
strategy VARCHAR(100) NOT NULL DEFAULT 'fifo',
|
||||
in_flow INTEGER NOT NULL DEFAULT 0,
|
||||
out_flow INTEGER NOT NULL DEFAULT 0,
|
||||
created_time INTEGER NOT NULL,
|
||||
updated_time INTEGER NOT NULL,
|
||||
status INTEGER NOT NULL,
|
||||
inx INTEGER NOT NULL DEFAULT 0,
|
||||
speed_id INTEGER
|
||||
)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("create legacy forward table: %v", err)
|
||||
}
|
||||
if err := db.Exec(`
|
||||
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx, speed_id)
|
||||
VALUES(1, 2, 'legacy-user', 'legacy-forward', 3, '127.0.0.1:9000', 'fifo', 0, 0, 1, 1, 1, 0, NULL)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert legacy forward: %v", err)
|
||||
}
|
||||
if sqlDB, _ := db.DB(); sqlDB != nil {
|
||||
_ = sqlDB.Close()
|
||||
}
|
||||
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open migrated sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
m := r.DB().Migrator()
|
||||
for _, field := range []string{"MaxConn", "IPMaxConn", "IPSpeedID", "ProxyProtocol"} {
|
||||
if !m.HasColumn(&model.Forward{}, field) {
|
||||
t.Fatalf("expected forward.%s column to exist", field)
|
||||
}
|
||||
}
|
||||
|
||||
var maxConn, ipMaxConn, proxyProtocol int
|
||||
var ipSpeedID sql.NullInt64
|
||||
if err := r.DB().Raw(`SELECT max_conn, ip_max_conn, ip_speed_id, proxy_protocol FROM forward WHERE id = 1`).Row().Scan(&maxConn, &ipMaxConn, &ipSpeedID, &proxyProtocol); err != nil {
|
||||
t.Fatalf("query forward defaults: %v", err)
|
||||
}
|
||||
if maxConn != 0 || ipMaxConn != 0 || ipSpeedID.Valid || proxyProtocol != 0 {
|
||||
t.Fatalf("expected default forward columns 0/0/NULL/0, got max_conn=%d ip_max_conn=%d ip_speed_id=%+v proxy_protocol=%d", maxConn, ipMaxConn, ipSpeedID, proxyProtocol)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMigrateSchemaRunsPostgresIDRepairEvenAtCurrentVersion(t *testing.T) {
|
||||
db, err := gorm.Open(gsqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
|
||||
@@ -397,22 +397,24 @@ 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, protocol 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, probeTargetHost string, probeTargetPort int, now int64) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
return tx.Model(&model.Tunnel{}).
|
||||
Where("id = ?", tunnelID).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"type": typeVal,
|
||||
"flow": flow,
|
||||
"traffic_ratio": trafficRatio,
|
||||
"status": status,
|
||||
"in_ip": nullStringFromInterface(inIP),
|
||||
"ip_preference": ipPreference,
|
||||
"protocol": protocol,
|
||||
"updated_time": now,
|
||||
"name": name,
|
||||
"type": typeVal,
|
||||
"flow": flow,
|
||||
"traffic_ratio": trafficRatio,
|
||||
"status": status,
|
||||
"in_ip": nullStringFromInterface(inIP),
|
||||
"ip_preference": ipPreference,
|
||||
"protocol": protocol,
|
||||
"probe_target_host": probeTargetHost,
|
||||
"probe_target_port": probeTargetPort,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -695,20 +697,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{}, maxConn int) error {
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) 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),
|
||||
"max_conn": maxConn,
|
||||
"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
|
||||
}
|
||||
|
||||
@@ -782,23 +787,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{}, maxConn int, now int64) {
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, 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),
|
||||
"max_conn": maxConn,
|
||||
"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
|
||||
}
|
||||
|
||||
@@ -1258,27 +1266,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{}, maxConn int) (int64, error) {
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) (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,
|
||||
MaxConn: maxConn,
|
||||
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
|
||||
@@ -1317,20 +1328,22 @@ func (r *Repository) BatchUpdateForwardStatus(ids []int64, status int) (int, int
|
||||
return s, f
|
||||
}
|
||||
|
||||
func (r *Repository) CreateTunnelTx(tx *gorm.DB, name string, trafficRatio float64, typeVal int, flow int64, now int64, status int, inIP interface{}, inx int, ipPreference string) (int64, error) {
|
||||
func (r *Repository) CreateTunnelTx(tx *gorm.DB, name string, trafficRatio float64, typeVal int, flow int64, now int64, status int, inIP interface{}, inx int, ipPreference string, probeTargetHost string, probeTargetPort int) (int64, error) {
|
||||
inIPVal := nullStringFromInterface(inIP)
|
||||
tunnel := model.Tunnel{
|
||||
Name: name,
|
||||
TrafficRatio: trafficRatio,
|
||||
Type: typeVal,
|
||||
Protocol: "tls",
|
||||
Flow: flow,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: status,
|
||||
InIP: inIPVal,
|
||||
Inx: inx,
|
||||
IPPreference: ipPreference,
|
||||
Name: name,
|
||||
TrafficRatio: trafficRatio,
|
||||
Type: typeVal,
|
||||
Protocol: "tls",
|
||||
Flow: flow,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: status,
|
||||
InIP: inIPVal,
|
||||
Inx: inx,
|
||||
IPPreference: ipPreference,
|
||||
ProbeTargetHost: probeTargetHost,
|
||||
ProbeTargetPort: probeTargetPort,
|
||||
}
|
||||
if err := tx.Create(&tunnel).Error; err != nil {
|
||||
return 0, 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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
type DatabaseStorageSummary struct {
|
||||
DBType string `json:"dbType"`
|
||||
DatabaseSizeBytes int64 `json:"databaseSizeBytes"`
|
||||
DatabaseSizeText string `json:"databaseSizeText"`
|
||||
}
|
||||
|
||||
func (r *Repository) DatabaseStorageSummary() (DatabaseStorageSummary, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return DatabaseStorageSummary{}, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
switch r.db.Dialector.Name() {
|
||||
case "sqlite":
|
||||
size, err := sqliteDatabaseFileSize(r.dbPath)
|
||||
if err != nil {
|
||||
return DatabaseStorageSummary{}, err
|
||||
}
|
||||
return DatabaseStorageSummary{DBType: "sqlite", DatabaseSizeBytes: size, DatabaseSizeText: formatDatabaseSize(size)}, nil
|
||||
case "postgres":
|
||||
var size int64
|
||||
if err := r.db.Raw("SELECT pg_database_size(current_database())").Scan(&size).Error; err != nil {
|
||||
return DatabaseStorageSummary{}, err
|
||||
}
|
||||
return DatabaseStorageSummary{DBType: "postgres", DatabaseSizeBytes: size, DatabaseSizeText: formatDatabaseSize(size)}, nil
|
||||
default:
|
||||
return DatabaseStorageSummary{}, fmt.Errorf("unsupported database dialect %q", r.db.Dialector.Name())
|
||||
}
|
||||
}
|
||||
|
||||
func sqliteDatabaseFileSize(path string) (int64, error) {
|
||||
if path == "" || path == ":memory:" {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
var total int64
|
||||
for _, candidate := range []string{path, path + "-wal", path + "-shm"} {
|
||||
info, err := os.Stat(candidate)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
continue
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
if !info.IsDir() {
|
||||
total += info.Size()
|
||||
}
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func formatDatabaseSize(bytes int64) string {
|
||||
if bytes < 1024 {
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
units := []string{"KB", "MB", "GB", "TB"}
|
||||
value := float64(bytes) / 1024
|
||||
for _, unit := range units {
|
||||
if value < 1024 || unit == "TB" {
|
||||
return fmt.Sprintf("%.1f %s", value, unit)
|
||||
}
|
||||
value /= 1024
|
||||
}
|
||||
return fmt.Sprintf("%d B", bytes)
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestDatabaseStorageSummarySQLiteIncludesSize(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "storage.db")
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
if err := r.InsertNodeMetric(&model.NodeMetric{NodeID: 1, Timestamp: 123, CPUUsage: 1}); err != nil {
|
||||
t.Fatalf("insert metric: %v", err)
|
||||
}
|
||||
|
||||
summary, err := r.DatabaseStorageSummary()
|
||||
if err != nil {
|
||||
t.Fatalf("storage summary: %v", err)
|
||||
}
|
||||
if summary.DBType != "sqlite" {
|
||||
t.Fatalf("expected sqlite db type, got %q", summary.DBType)
|
||||
}
|
||||
if summary.DatabaseSizeBytes <= 0 {
|
||||
t.Fatalf("expected database size > 0, got %d", summary.DatabaseSizeBytes)
|
||||
}
|
||||
if summary.DatabaseSizeText == "" {
|
||||
t.Fatalf("expected formatted size")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatDatabaseSize(t *testing.T) {
|
||||
tests := []struct {
|
||||
bytes int64
|
||||
want string
|
||||
}{
|
||||
{bytes: 0, want: "0 B"},
|
||||
{bytes: 512, want: "512 B"},
|
||||
{bytes: 1024, want: "1.0 KB"},
|
||||
{bytes: 1024 * 1024, want: "1.0 MB"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
if got := formatDatabaseSize(tc.bytes); got != tc.want {
|
||||
t.Fatalf("formatDatabaseSize(%d) = %q, want %q", tc.bytes, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -57,13 +58,14 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
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)
|
||||
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)
|
||||
`, tunnelID, now+365*24*3600*1000).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
var commandMu sync.Mutex
|
||||
@@ -91,11 +93,13 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
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,
|
||||
"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 {
|
||||
@@ -120,6 +124,47 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
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()
|
||||
|
||||
@@ -152,8 +197,8 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
t.Fatalf("expected limiter name %s, got %v", expectedName, addData["name"])
|
||||
}
|
||||
if limits, ok := addData["limits"].([]interface{}); ok {
|
||||
if len(limits) != 1 || limits[0] != "$ 42" {
|
||||
t.Fatalf("expected limits to contain '$ 42', got %v", limits)
|
||||
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)
|
||||
@@ -175,14 +220,290 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
t.Fatalf("expected nested name %s, got %v", expectedName, nestedData["name"])
|
||||
}
|
||||
if nestedLimits, ok := nestedData["limits"].([]interface{}); ok {
|
||||
if len(nestedLimits) != 1 || nestedLimits[0] != "$ 42" {
|
||||
t.Fatalf("expected nested limits to contain '$ 42', got %v", nestedLimits)
|
||||
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()
|
||||
|
||||
|
||||
@@ -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,64 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestStorageSummaryRequiresAdminAndReturnsSize(t *testing.T) {
|
||||
secret := "storage-contract-secret"
|
||||
router, _ := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
userToken, err := auth.GenerateToken(2, "normal_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate user token: %v", err)
|
||||
}
|
||||
|
||||
userReq := httptest.NewRequest(http.MethodGet, "/api/v1/system/storage", nil)
|
||||
userReq.Header.Set("Authorization", userToken)
|
||||
userRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(userRes, userReq)
|
||||
|
||||
var denied response.R
|
||||
if err := json.NewDecoder(userRes.Body).Decode(&denied); err != nil {
|
||||
t.Fatalf("decode denied response: %v", err)
|
||||
}
|
||||
if denied.Code != 403 {
|
||||
t.Fatalf("expected 403 for non-admin, got %d", denied.Code)
|
||||
}
|
||||
|
||||
adminReq := httptest.NewRequest(http.MethodGet, "/api/v1/system/storage", nil)
|
||||
adminReq.Header.Set("Authorization", adminToken)
|
||||
adminRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(adminRes, adminReq)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(adminRes.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode admin response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d: %s", out.Code, out.Msg)
|
||||
}
|
||||
data, ok := out.Data.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected object data, got %T", out.Data)
|
||||
}
|
||||
if data["dbType"] == "" {
|
||||
t.Fatalf("expected dbType")
|
||||
}
|
||||
if _, ok := data["databaseSizeBytes"].(float64); !ok {
|
||||
t.Fatalf("expected numeric databaseSizeBytes, got %T", data["databaseSizeBytes"])
|
||||
}
|
||||
if data["databaseSizeText"] == "" {
|
||||
t.Fatalf("expected databaseSizeText")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -229,6 +234,13 @@ func (p *program) reloadConfig() error {
|
||||
if err := loader.Load(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
activeServices := make(map[string]struct{}, len(cfg.Services))
|
||||
for _, svc := range cfg.Services {
|
||||
if svc != nil {
|
||||
activeServices[svc.Name] = struct{}{}
|
||||
}
|
||||
}
|
||||
xservice.GetGlobalTrafficManager().RetainServices(activeServices)
|
||||
|
||||
if err := p.run(cfg); err != nil {
|
||||
return err
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/go-gost/x/config/loader"
|
||||
"github.com/go-gost/x/config/parsing/parser"
|
||||
"github.com/go-gost/x/registry"
|
||||
xservice "github.com/go-gost/x/service"
|
||||
)
|
||||
|
||||
// swagger:parameters reloadConfigRequest
|
||||
@@ -42,6 +43,13 @@ func reloadConfig(ctx *gin.Context) {
|
||||
writeError(ctx, NewError(http.StatusBadRequest, ErrCodeInvalid, err.Error()))
|
||||
return
|
||||
}
|
||||
activeServices := make(map[string]struct{}, len(cfg.Services))
|
||||
for _, svc := range cfg.Services {
|
||||
if svc != nil {
|
||||
activeServices[svc.Name] = struct{}{}
|
||||
}
|
||||
}
|
||||
xservice.GetGlobalTrafficManager().RetainServices(activeServices)
|
||||
|
||||
for _, svc := range registry.ServiceRegistry().GetAll() {
|
||||
svc := svc
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
parser "github.com/go-gost/x/config/parsing/service"
|
||||
kill "github.com/go-gost/x/internal/util/port"
|
||||
"github.com/go-gost/x/registry"
|
||||
xservice "github.com/go-gost/x/service"
|
||||
)
|
||||
|
||||
// swagger:parameters createServiceRequest
|
||||
@@ -409,6 +410,7 @@ func deleteService(ctx *gin.Context) {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
xservice.GetGlobalTrafficManager().RemoveServices(name)
|
||||
|
||||
ctx.JSON(http.StatusOK, Response{
|
||||
Msg: "OK",
|
||||
@@ -484,6 +486,12 @@ func deleteServices(ctx *gin.Context) {
|
||||
return nil
|
||||
})
|
||||
|
||||
names := make([]string, 0, len(servicesToDelete))
|
||||
for _, std := range servicesToDelete {
|
||||
names = append(names, std.name)
|
||||
}
|
||||
xservice.GetGlobalTrafficManager().RemoveServices(names...)
|
||||
|
||||
ctx.JSON(http.StatusOK, Response{
|
||||
Msg: "OK",
|
||||
})
|
||||
|
||||
@@ -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 {
|
||||
err = 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")
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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),
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
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() error {
|
||||
persistMu.Lock()
|
||||
path := persistPath
|
||||
enabled := persistEnable
|
||||
persistMu.Unlock()
|
||||
|
||||
if !enabled || path == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
cfg := Global()
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
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 fmt.Errorf("config persist: marshal failed: %w", err)
|
||||
}
|
||||
|
||||
// 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 fmt.Errorf("config persist: create temp file failed: %w", err)
|
||||
}
|
||||
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 fmt.Errorf("config persist: write failed: %w", err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
os.Remove(tmpName)
|
||||
fmt.Printf("⚠️ config persist: close temp file failed: %v\n", err)
|
||||
return fmt.Errorf("config persist: close temp file failed: %w", err)
|
||||
}
|
||||
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
os.Remove(tmpName)
|
||||
fmt.Printf("⚠️ config persist: rename failed: %v\n", err)
|
||||
return fmt.Errorf("config persist: rename failed: %w", err)
|
||||
}
|
||||
|
||||
fmt.Printf("💾 节点配置已持久化到 %s\n", path)
|
||||
return nil
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -12,10 +12,25 @@ type chainRegistry struct {
|
||||
registry[chain.Chainer]
|
||||
}
|
||||
|
||||
func ReplaceChain(name string, v chain.Chainer) error {
|
||||
if name == "" {
|
||||
return nil
|
||||
}
|
||||
if r, ok := chainReg.(*chainRegistry); ok {
|
||||
r.replace(name, v)
|
||||
return nil
|
||||
}
|
||||
return chainReg.Register(name, v)
|
||||
}
|
||||
|
||||
func (r *chainRegistry) Register(name string, v chain.Chainer) error {
|
||||
return r.registry.Register(name, v)
|
||||
}
|
||||
|
||||
func (r *chainRegistry) replace(name string, v chain.Chainer) {
|
||||
r.m.Store(name, v)
|
||||
}
|
||||
|
||||
func (r *chainRegistry) Get(name string) chain.Chainer {
|
||||
if name != "" {
|
||||
return &chainWrapper{name: name, r: r}
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package registry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
)
|
||||
|
||||
type testChainer struct {
|
||||
route chain.Route
|
||||
}
|
||||
|
||||
func (c testChainer) Route(context.Context, string, string, ...chain.RouteOption) chain.Route {
|
||||
return c.route
|
||||
}
|
||||
|
||||
type testRoute struct {
|
||||
nodes []*chain.Node
|
||||
}
|
||||
|
||||
func (r testRoute) Dial(context.Context, string, string, ...chain.DialOption) (net.Conn, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (r testRoute) Bind(context.Context, string, string, ...chain.BindOption) (net.Listener, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (r testRoute) Nodes() []*chain.Node {
|
||||
return r.nodes
|
||||
}
|
||||
|
||||
func TestReplaceChainOverwritesExistingRegistration(t *testing.T) {
|
||||
name := "replace_chain_tdd"
|
||||
ChainRegistry().Unregister(name)
|
||||
defer ChainRegistry().Unregister(name)
|
||||
|
||||
if err := ChainRegistry().Register(name, testChainer{route: testRoute{nodes: []*chain.Node{{Name: "old"}}}}); err != nil {
|
||||
t.Fatalf("register old chain: %v", err)
|
||||
}
|
||||
if err := ReplaceChain(name, testChainer{route: testRoute{nodes: []*chain.Node{{Name: "new"}}}}); err != nil {
|
||||
t.Fatalf("replace chain: %v", err)
|
||||
}
|
||||
|
||||
route := ChainRegistry().Get(name).Route(context.Background(), "tcp", "example.com:443")
|
||||
if route == nil || len(route.Nodes()) != 1 || route.Nodes()[0].Name != "new" {
|
||||
t.Fatalf("expected replacement chain route, got %#v", route)
|
||||
}
|
||||
}
|
||||
@@ -5,15 +5,17 @@ import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
// GlobalTrafficManager 全局流量管理器(所有服务共享)
|
||||
type GlobalTrafficManager struct {
|
||||
mu sync.RWMutex
|
||||
mu sync.RWMutex
|
||||
serviceTraffic map[string]*ServiceTraffic // key: 服务名, value: 流量数据
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
reportTicker *time.Ticker
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
reportTicker *time.Ticker
|
||||
}
|
||||
|
||||
// ServiceTraffic 单个服务的流量累积
|
||||
@@ -50,6 +52,9 @@ func (m *GlobalTrafficManager) AddTraffic(serviceName string, upBytes, downBytes
|
||||
if upBytes == 0 && downBytes == 0 {
|
||||
return
|
||||
}
|
||||
if !registry.ServiceRegistry().IsRegistered(serviceName) {
|
||||
return
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
@@ -70,6 +75,51 @@ func (m *GlobalTrafficManager) AddTraffic(serviceName string, upBytes, downBytes
|
||||
traffic.mu.Unlock()
|
||||
}
|
||||
|
||||
// RemoveServices drops cached traffic counters for services that no longer exist.
|
||||
func (m *GlobalTrafficManager) RemoveServices(serviceNames ...string) {
|
||||
if m == nil || len(serviceNames) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for _, name := range serviceNames {
|
||||
if m.isTrafficEmptyLocked(name) {
|
||||
delete(m.serviceTraffic, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RetainServices removes traffic counters for every service not in activeNames.
|
||||
func (m *GlobalTrafficManager) RetainServices(activeNames map[string]struct{}) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
for name := range m.serviceTraffic {
|
||||
if _, ok := activeNames[name]; !ok {
|
||||
if m.isTrafficEmptyLocked(name) {
|
||||
delete(m.serviceTraffic, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *GlobalTrafficManager) isTrafficEmptyLocked(name string) bool {
|
||||
traffic, ok := m.serviceTraffic[name]
|
||||
if !ok {
|
||||
return true
|
||||
}
|
||||
|
||||
traffic.mu.Lock()
|
||||
defer traffic.mu.Unlock()
|
||||
return traffic.UpBytes == 0 && traffic.DownBytes == 0
|
||||
}
|
||||
|
||||
// startReporting 启动定时上报协程(每5秒执行一次)
|
||||
func (m *GlobalTrafficManager) startReporting() {
|
||||
|
||||
@@ -105,6 +155,7 @@ func (m *GlobalTrafficManager) collectAndReport() {
|
||||
traffic.DownBytes = 0
|
||||
}
|
||||
traffic.mu.Unlock()
|
||||
isStale := !registry.ServiceRegistry().IsRegistered(name)
|
||||
|
||||
if up > 0 || down > 0 {
|
||||
reportItems = append(reportItems, TrafficReportItem{
|
||||
@@ -113,6 +164,9 @@ func (m *GlobalTrafficManager) collectAndReport() {
|
||||
D: down,
|
||||
})
|
||||
}
|
||||
if isStale {
|
||||
delete(m.serviceTraffic, name)
|
||||
}
|
||||
}
|
||||
|
||||
m.mu.Unlock()
|
||||
@@ -162,4 +216,3 @@ func (m *GlobalTrafficManager) GetServiceTraffic(serviceName string) (upBytes, d
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestGlobalTrafficManagerRemoveServicesDropsCachedEntries(t *testing.T) {
|
||||
m := &GlobalTrafficManager{serviceTraffic: map[string]*ServiceTraffic{
|
||||
"svc-a": {ServiceName: "svc-a"},
|
||||
"svc-b": {ServiceName: "svc-b"},
|
||||
}}
|
||||
|
||||
m.RemoveServices("svc-a")
|
||||
|
||||
if _, ok := m.serviceTraffic["svc-a"]; ok {
|
||||
t.Fatalf("expected svc-a traffic entry to be removed")
|
||||
}
|
||||
if _, ok := m.serviceTraffic["svc-b"]; !ok {
|
||||
t.Fatalf("expected svc-b traffic entry to remain")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGlobalTrafficManagerRetainServicesDropsStaleEntries(t *testing.T) {
|
||||
m := &GlobalTrafficManager{serviceTraffic: map[string]*ServiceTraffic{
|
||||
"svc-a": {ServiceName: "svc-a"},
|
||||
"svc-b": {ServiceName: "svc-b"},
|
||||
}}
|
||||
|
||||
m.RetainServices(map[string]struct{}{"svc-b": {}})
|
||||
|
||||
if _, ok := m.serviceTraffic["svc-a"]; ok {
|
||||
t.Fatalf("expected stale svc-a traffic entry to be removed")
|
||||
}
|
||||
if _, ok := m.serviceTraffic["svc-b"]; !ok {
|
||||
t.Fatalf("expected active svc-b traffic entry to remain")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGlobalTrafficManagerAddTrafficIgnoresUnregisteredService(t *testing.T) {
|
||||
m := &GlobalTrafficManager{serviceTraffic: make(map[string]*ServiceTraffic)}
|
||||
|
||||
m.AddTraffic("deleted-service", 10, 20)
|
||||
|
||||
if _, ok := m.serviceTraffic["deleted-service"]; ok {
|
||||
t.Fatalf("expected unregistered service traffic to be ignored")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGlobalTrafficManagerCollectAndReportDropsStaleEntriesAfterReporting(t *testing.T) {
|
||||
origReportDo := reportDo
|
||||
origReportURL := httpReportURL
|
||||
origAESCrypto := httpAESCrypto
|
||||
defer func() {
|
||||
reportDo = origReportDo
|
||||
httpReportURL = origReportURL
|
||||
httpAESCrypto = origAESCrypto
|
||||
}()
|
||||
|
||||
httpReportURL = "http://panel.example.com/flow/upload?secret=abc"
|
||||
httpAESCrypto = nil
|
||||
|
||||
var requestBody string
|
||||
reportDo = func(_ context.Context, req *http.Request, _ time.Duration) (*http.Response, error) {
|
||||
body, err := io.ReadAll(req.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("read request body: %v", err)
|
||||
}
|
||||
requestBody = string(body)
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader("ok")),
|
||||
}, nil
|
||||
}
|
||||
|
||||
m := &GlobalTrafficManager{
|
||||
serviceTraffic: map[string]*ServiceTraffic{
|
||||
"stale-service": {ServiceName: "stale-service", UpBytes: 10, DownBytes: 20},
|
||||
},
|
||||
ctx: context.Background(),
|
||||
}
|
||||
|
||||
m.collectAndReport()
|
||||
|
||||
if !strings.Contains(requestBody, "stale-service") {
|
||||
t.Fatalf("expected pending stale traffic to be reported first, body=%s", requestBody)
|
||||
}
|
||||
if _, ok := m.serviceTraffic["stale-service"]; ok {
|
||||
t.Fatalf("expected stale traffic entry to be removed after report collection")
|
||||
}
|
||||
}
|
||||
@@ -31,34 +31,32 @@ func createChain(req createChainRequest) error {
|
||||
return errors.New("chain " + name + " already exists")
|
||||
}
|
||||
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
c.Chains = append(c.Chains, &req.Data)
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateChain(req updateChainRequest) error {
|
||||
|
||||
name := strings.TrimSpace(req.Chain)
|
||||
|
||||
if registry.ChainRegistry().IsRegistered(name) {
|
||||
registry.ChainRegistry().Unregister(name)
|
||||
if name == "" {
|
||||
name = strings.TrimSpace(req.Data.Name)
|
||||
}
|
||||
if name == "" {
|
||||
return errors.New("chain name is required")
|
||||
}
|
||||
|
||||
req.Data.Name = name
|
||||
|
||||
v, err := parser.ParseChain(&req.Data, logger.Default())
|
||||
if err != nil {
|
||||
return errors.New("create chain " + name + " failed: " + err.Error())
|
||||
}
|
||||
|
||||
if err := registry.ChainRegistry().Register(name, v); err != nil {
|
||||
if err := registry.ReplaceChain(name, v); err != nil {
|
||||
return errors.New("chain " + name + " already exists")
|
||||
}
|
||||
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
found := false
|
||||
for i := range c.Chains {
|
||||
if c.Chains[i].Name == name {
|
||||
@@ -72,8 +70,6 @@ func updateChain(req updateChainRequest) error {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteChain(req deleteChainRequest) error {
|
||||
@@ -84,7 +80,7 @@ func deleteChain(req deleteChainRequest) error {
|
||||
registry.ChainRegistry().Unregister(name)
|
||||
}
|
||||
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
chains := c.Chains
|
||||
c.Chains = nil
|
||||
for _, s := range chains {
|
||||
@@ -95,8 +91,6 @@ func deleteChain(req deleteChainRequest) error {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type createChainRequest struct {
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
corelogger "github.com/go-gost/core/logger"
|
||||
"github.com/go-gost/x/config"
|
||||
_ "github.com/go-gost/x/connector/relay"
|
||||
_ "github.com/go-gost/x/dialer/tcp"
|
||||
xlogger "github.com/go-gost/x/logger"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
func TestUpdateChainParseFailureKeepsExistingChainRegistered(t *testing.T) {
|
||||
corelogger.SetDefault(xlogger.Nop())
|
||||
|
||||
name := "chain_update_parse_failure_tdd"
|
||||
originalConfig := config.Global()
|
||||
defer config.Set(originalConfig)
|
||||
registry.ChainRegistry().Unregister(name)
|
||||
defer registry.ChainRegistry().Unregister(name)
|
||||
config.Set(&config.Config{})
|
||||
|
||||
valid := config.ChainConfig{
|
||||
Name: name,
|
||||
Hops: []*config.HopConfig{{
|
||||
Name: "hop-valid",
|
||||
Nodes: []*config.NodeConfig{{
|
||||
Name: "node-valid",
|
||||
Addr: "127.0.0.1:443",
|
||||
Connector: &config.ConnectorConfig{Type: "relay"},
|
||||
Dialer: &config.DialerConfig{Type: "tcp"},
|
||||
}},
|
||||
}},
|
||||
}
|
||||
if err := createChain(createChainRequest{Data: valid}); err != nil {
|
||||
t.Fatalf("create valid chain: %v", err)
|
||||
}
|
||||
before := registry.ChainRegistry().Get(name)
|
||||
if before == nil || !registry.ChainRegistry().IsRegistered(name) {
|
||||
t.Fatalf("expected chain registered before update")
|
||||
}
|
||||
|
||||
invalid := config.ChainConfig{
|
||||
Hops: []*config.HopConfig{{
|
||||
Name: "hop-invalid",
|
||||
Nodes: []*config.NodeConfig{{
|
||||
Name: "node-invalid",
|
||||
Addr: "127.0.0.1:443",
|
||||
Connector: &config.ConnectorConfig{Type: "connector-does-not-exist"},
|
||||
Dialer: &config.DialerConfig{Type: "tcp"},
|
||||
}},
|
||||
}},
|
||||
}
|
||||
err := updateChain(updateChainRequest{Chain: name, Data: invalid})
|
||||
if err == nil {
|
||||
t.Fatalf("expected invalid chain update to fail")
|
||||
}
|
||||
if !registry.ChainRegistry().IsRegistered(name) {
|
||||
t.Fatalf("expected old chain to remain registered after failed update")
|
||||
}
|
||||
cfg := config.Global()
|
||||
if len(cfg.Chains) != 1 || cfg.Chains[0] == nil || cfg.Chains[0].Name != name {
|
||||
t.Fatalf("expected original chain config to remain, got %#v", cfg.Chains)
|
||||
}
|
||||
if got := cfg.Chains[0].Hops[0].Name; got != "hop-valid" {
|
||||
t.Fatalf("expected original chain config to remain, got hop %q", got)
|
||||
}
|
||||
}
|
||||
+12
-18
@@ -25,12 +25,10 @@ func createLimiter(req createLimiterRequest) error {
|
||||
return errors.New("limiter " + name + " already exists")
|
||||
}
|
||||
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
c.Limiters = append(c.Limiters, &req.Data)
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateLimiter(req updateLimiterRequest) error {
|
||||
@@ -49,7 +47,7 @@ func updateLimiter(req updateLimiterRequest) error {
|
||||
return errors.New("limiter " + name + " already exists")
|
||||
}
|
||||
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
found := false
|
||||
for i := range c.Limiters {
|
||||
if c.Limiters[i].Name == name {
|
||||
@@ -63,8 +61,6 @@ func updateLimiter(req updateLimiterRequest) error {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func deleteLimiter(req deleteLimiterRequest) error {
|
||||
@@ -75,7 +71,7 @@ func deleteLimiter(req deleteLimiterRequest) error {
|
||||
registry.TrafficLimiterRegistry().Unregister(name)
|
||||
}
|
||||
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
limiteres := c.Limiters
|
||||
c.Limiters = nil
|
||||
for _, s := range limiteres {
|
||||
@@ -86,8 +82,6 @@ func deleteLimiter(req deleteLimiterRequest) error {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type createLimiterRequest struct {
|
||||
@@ -120,10 +114,10 @@ func createConnLimiter(req createLimiterRequest) error {
|
||||
return errors.New("conn limiter " + name + " already exists")
|
||||
}
|
||||
|
||||
if c := config.Global(); c != nil {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
c.CLimiters = append(c.CLimiters, &req.Data)
|
||||
}
|
||||
return nil
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func updateConnLimiter(req updateLimiterRequest) error {
|
||||
@@ -139,7 +133,7 @@ func updateConnLimiter(req updateLimiterRequest) error {
|
||||
return errors.New("conn limiter " + name + " already exists")
|
||||
}
|
||||
|
||||
if c := config.Global(); c != nil {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
for i := range c.CLimiters {
|
||||
if c.CLimiters[i].Name == name {
|
||||
c.CLimiters[i] = &req.Data
|
||||
@@ -147,8 +141,8 @@ func updateConnLimiter(req updateLimiterRequest) error {
|
||||
}
|
||||
}
|
||||
c.CLimiters = append(c.CLimiters, &req.Data)
|
||||
}
|
||||
return nil
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func deleteConnLimiter(req deleteLimiterRequest) error {
|
||||
@@ -158,7 +152,7 @@ func deleteConnLimiter(req deleteLimiterRequest) error {
|
||||
registry.ConnLimiterRegistry().Unregister(name)
|
||||
}
|
||||
|
||||
if c := config.Global(); c != nil {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
limiteres := c.CLimiters
|
||||
c.CLimiters = nil
|
||||
for _, s := range limiteres {
|
||||
@@ -167,6 +161,6 @@ func deleteConnLimiter(req deleteLimiterRequest) error {
|
||||
}
|
||||
c.CLimiters = append(c.CLimiters, s)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
corelogger "github.com/go-gost/core/logger"
|
||||
"github.com/go-gost/x/config"
|
||||
xlogger "github.com/go-gost/x/logger"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
func TestCreateConnLimiterUpdatesGlobalConfig(t *testing.T) {
|
||||
corelogger.SetDefault(xlogger.Nop())
|
||||
|
||||
name := "conn_limiter_tdd"
|
||||
originalConfig := config.Global()
|
||||
defer config.Set(originalConfig)
|
||||
registry.ConnLimiterRegistry().Unregister(name)
|
||||
defer registry.ConnLimiterRegistry().Unregister(name)
|
||||
config.Set(&config.Config{})
|
||||
|
||||
err := createConnLimiter(createLimiterRequest{Data: config.LimiterConfig{Name: name, Limits: []string{"$ 1"}}})
|
||||
if err != nil {
|
||||
t.Fatalf("create conn limiter: %v", err)
|
||||
}
|
||||
|
||||
cfg := config.Global()
|
||||
if len(cfg.CLimiters) != 1 || cfg.CLimiters[0] == nil || cfg.CLimiters[0].Name != name {
|
||||
t.Fatalf("expected conn limiter in global config, got %#v", cfg.CLimiters)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateLimiterReportsPersistFailure(t *testing.T) {
|
||||
corelogger.SetDefault(xlogger.Nop())
|
||||
|
||||
name := "traffic_limiter_persist_tdd"
|
||||
originalConfig := config.Global()
|
||||
originalPersistPath := config.PersistPath()
|
||||
defer config.Set(originalConfig)
|
||||
defer config.SetPersistPath(originalPersistPath)
|
||||
registry.TrafficLimiterRegistry().Unregister(name)
|
||||
defer registry.TrafficLimiterRegistry().Unregister(name)
|
||||
config.Set(&config.Config{})
|
||||
config.SetPersistPath(filepath.Join(t.TempDir(), "missing", "gost.json"))
|
||||
config.EnablePersist()
|
||||
|
||||
err := createLimiter(createLimiterRequest{Data: config.LimiterConfig{Name: name, Limits: []string{"$ 1"}}})
|
||||
if err == nil {
|
||||
t.Fatalf("expected persist failure to be returned")
|
||||
}
|
||||
}
|
||||
+45
-15
@@ -3,6 +3,7 @@ package socket
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
parser "github.com/go-gost/x/config/parsing/service"
|
||||
kill "github.com/go-gost/x/internal/util/port"
|
||||
"github.com/go-gost/x/registry"
|
||||
xservice "github.com/go-gost/x/service"
|
||||
)
|
||||
|
||||
func createServices(req createServicesRequest) error {
|
||||
@@ -53,9 +55,8 @@ func createServices(req createServicesRequest) error {
|
||||
if err := registry.ServiceRegistry().Register(ps.config.Name, ps.service); err != nil {
|
||||
// 如果注册失败,回滚已注册的服务
|
||||
for _, regName := range registeredServices {
|
||||
if svc := registry.ServiceRegistry().Get(regName); svc != nil {
|
||||
if registry.ServiceRegistry().Get(regName) != nil {
|
||||
registry.ServiceRegistry().Unregister(regName)
|
||||
svc.Close()
|
||||
}
|
||||
}
|
||||
return errors.New("service " + ps.config.Name + " already exists")
|
||||
@@ -71,14 +72,12 @@ func createServices(req createServicesRequest) error {
|
||||
}
|
||||
|
||||
// 第四阶段:更新配置
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
return config.OnUpdate(func(c *config.Config) error {
|
||||
for _, ps := range parsedServices {
|
||||
c.Services = append(c.Services, &ps.config)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func updateServices(req updateServicesRequest) error {
|
||||
@@ -97,17 +96,23 @@ func updateServices(req updateServicesRequest) error {
|
||||
}
|
||||
|
||||
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)
|
||||
changedServices := make([]struct {
|
||||
config config.ServiceConfig
|
||||
service service.Service
|
||||
}, 0, len(req.Data))
|
||||
for i := range req.Data {
|
||||
serviceConfig := &req.Data[i]
|
||||
name := serviceConfig.Name
|
||||
if registry.ServiceRegistry().Get(name) != nil && serviceConfigUnchanged(name, *serviceConfig) {
|
||||
continue
|
||||
}
|
||||
|
||||
// 1. 获取旧服务
|
||||
old := registry.ServiceRegistry().Get(name)
|
||||
|
||||
// 2. 关闭旧服务 (如果存在)
|
||||
if old != nil {
|
||||
old.Close()
|
||||
// 3. 从注册表移除旧服务
|
||||
// 3. 从注册表移除旧服务;registry 会负责关闭旧服务。
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
|
||||
@@ -116,6 +121,10 @@ func updateServices(req updateServicesRequest) error {
|
||||
if err != nil {
|
||||
return errors.New("create service " + name + " failed: " + err.Error())
|
||||
}
|
||||
changedServices = append(changedServices, struct {
|
||||
config config.ServiceConfig
|
||||
service service.Service
|
||||
}{*serviceConfig, svc})
|
||||
|
||||
// 5. 注册新服务
|
||||
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
||||
@@ -126,12 +135,15 @@ func updateServices(req updateServicesRequest) error {
|
||||
// 6. 启动新服务
|
||||
go svc.Serve()
|
||||
}
|
||||
if len(changedServices) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 第三阶段:更新配置
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
for i := range req.Data {
|
||||
if err := config.OnUpdate(func(c *config.Config) error {
|
||||
for i := range changedServices {
|
||||
// 创建副本以确保指针安全
|
||||
cfgCopy := req.Data[i]
|
||||
cfgCopy := changedServices[i].config
|
||||
found := false
|
||||
for j := range c.Services {
|
||||
if c.Services[j].Name == cfgCopy.Name {
|
||||
@@ -145,11 +157,30 @@ func updateServices(req updateServicesRequest) error {
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
|
||||
cfg := config.Global()
|
||||
if cfg == nil {
|
||||
return false
|
||||
}
|
||||
next.Status = nil
|
||||
for _, current := range cfg.Services {
|
||||
if current == nil || strings.TrimSpace(current.Name) != name {
|
||||
continue
|
||||
}
|
||||
currentCopy := *current
|
||||
currentCopy.Status = nil
|
||||
return reflect.DeepEqual(currentCopy, next)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func deleteServices(req deleteServicesRequest) error {
|
||||
|
||||
if len(req.Services) == 0 {
|
||||
@@ -182,7 +213,6 @@ func deleteServices(req deleteServicesRequest) error {
|
||||
// 第二阶段:删除所有服务
|
||||
for _, std := range servicesToDelete {
|
||||
registry.ServiceRegistry().Unregister(std.name)
|
||||
std.service.Close()
|
||||
}
|
||||
// 确保所有请求删除的服务都从注册表中移除(即使之前未找到实例)
|
||||
for _, name := range namesToRemove {
|
||||
@@ -192,7 +222,7 @@ func deleteServices(req deleteServicesRequest) error {
|
||||
}
|
||||
|
||||
// 第三阶段:更新配置
|
||||
config.OnUpdate(func(c *config.Config) error {
|
||||
err := config.OnUpdate(func(c *config.Config) error {
|
||||
services := c.Services
|
||||
c.Services = nil
|
||||
for _, s := range services {
|
||||
@@ -209,8 +239,8 @@ func deleteServices(req deleteServicesRequest) error {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
return nil
|
||||
xservice.GetGlobalTrafficManager().RemoveServices(namesToRemove...)
|
||||
return err
|
||||
}
|
||||
|
||||
func pauseServices(req pauseServicesRequest) error {
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
corelogger "github.com/go-gost/core/logger"
|
||||
"github.com/go-gost/core/service"
|
||||
"github.com/go-gost/x/config"
|
||||
xlogger "github.com/go-gost/x/logger"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
type recordingService struct {
|
||||
closed int
|
||||
}
|
||||
|
||||
func (s *recordingService) Serve() error { return nil }
|
||||
func (s *recordingService) Addr() net.Addr { return nil }
|
||||
func (s *recordingService) Close() error {
|
||||
s.closed++
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestUpdateServicesSkipsUnchangedServiceWithoutRestart(t *testing.T) {
|
||||
corelogger.SetDefault(xlogger.Nop())
|
||||
|
||||
name := "unchanged_service_tdd"
|
||||
existing := &recordingService{}
|
||||
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
defer registry.ServiceRegistry().Unregister(name)
|
||||
if err := registry.ServiceRegistry().Register(name, service.Service(existing)); err != nil {
|
||||
t.Fatalf("register existing service: %v", err)
|
||||
}
|
||||
|
||||
originalConfig := config.Global()
|
||||
defer config.Set(originalConfig)
|
||||
serviceConfig := config.ServiceConfig{Name: name, Addr: "127.0.0.1:0"}
|
||||
config.Set(&config.Config{Services: []*config.ServiceConfig{&serviceConfig}})
|
||||
|
||||
if err := updateServices(updateServicesRequest{Data: []config.ServiceConfig{serviceConfig}}); err != nil {
|
||||
t.Fatalf("unchanged update should succeed without parsing/restarting: %v", err)
|
||||
}
|
||||
if existing.closed != 0 {
|
||||
t.Fatalf("unchanged service was restarted, closed %d times", existing.closed)
|
||||
}
|
||||
if got := registry.ServiceRegistry().Get(name); got != service.Service(existing) {
|
||||
t.Fatalf("expected existing service to remain registered")
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
@@ -157,6 +158,9 @@ type WebSocketReporter struct {
|
||||
addr string // 保存服务器地址
|
||||
secret string // 保存密钥
|
||||
version string // 保存版本号
|
||||
http int
|
||||
tls int
|
||||
socks int
|
||||
preferredWSScheme string
|
||||
conn *websocket.Conn
|
||||
curBackoff time.Duration // 当前重连退避间隔
|
||||
@@ -189,9 +193,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,
|
||||
@@ -295,9 +299,9 @@ func (w *WebSocketReporter) connect() error {
|
||||
Socks int `json:"socks"`
|
||||
}
|
||||
|
||||
var cfg LocalConfig
|
||||
cfg := LocalConfig{Http: w.http, Tls: w.tls, Socks: w.socks}
|
||||
if b, err := os.ReadFile("config.json"); err == nil {
|
||||
json.Unmarshal(b, &cfg)
|
||||
_ = json.Unmarshal(b, &cfg)
|
||||
}
|
||||
|
||||
candidates := buildWebSocketCandidates(w.addr, w.secret, w.version, cfg.Http, cfg.Tls, cfg.Socks, w.preferredWSScheme)
|
||||
@@ -781,7 +785,6 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
fmt.Println("🔔 收到命令: ", string(jsonBytes))
|
||||
var err error
|
||||
var response CommandResponse
|
||||
var needSaveConfig bool // 标记是否需要保存配置(只有状态变更命令才需要)
|
||||
|
||||
// 传递 requestId
|
||||
response.RequestId = cmd.RequestId
|
||||
@@ -791,63 +794,49 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
case "AddService":
|
||||
err = w.handleAddService(cmd.Data)
|
||||
response.Type = "AddServiceResponse"
|
||||
needSaveConfig = true
|
||||
case "UpdateService":
|
||||
err = w.handleUpdateService(cmd.Data)
|
||||
response.Type = "UpdateServiceResponse"
|
||||
needSaveConfig = true
|
||||
case "DeleteService":
|
||||
err = w.handleDeleteService(cmd.Data)
|
||||
response.Type = "DeleteServiceResponse"
|
||||
needSaveConfig = true
|
||||
case "PauseService":
|
||||
err = w.handlePauseService(cmd.Data)
|
||||
response.Type = "PauseServiceResponse"
|
||||
needSaveConfig = true
|
||||
case "ResumeService":
|
||||
err = w.handleResumeService(cmd.Data)
|
||||
response.Type = "ResumeServiceResponse"
|
||||
needSaveConfig = true
|
||||
|
||||
// Chain 相关命令
|
||||
case "AddChains":
|
||||
err = w.handleAddChain(cmd.Data)
|
||||
response.Type = "AddChainsResponse"
|
||||
needSaveConfig = true
|
||||
case "UpdateChains":
|
||||
err = w.handleUpdateChain(cmd.Data)
|
||||
response.Type = "UpdateChainsResponse"
|
||||
needSaveConfig = true
|
||||
case "DeleteChains":
|
||||
err = w.handleDeleteChain(cmd.Data)
|
||||
response.Type = "DeleteChainsResponse"
|
||||
needSaveConfig = true
|
||||
|
||||
// Limiter 相关命令
|
||||
case "AddLimiters":
|
||||
err = w.handleAddLimiter(cmd.Data)
|
||||
response.Type = "AddLimitersResponse"
|
||||
needSaveConfig = true
|
||||
case "UpdateLimiters":
|
||||
err = w.handleUpdateLimiter(cmd.Data)
|
||||
response.Type = "UpdateLimitersResponse"
|
||||
needSaveConfig = true
|
||||
case "DeleteLimiters":
|
||||
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":
|
||||
@@ -875,7 +864,6 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
case "SetProtocol":
|
||||
err = w.handleSetProtocol(cmd.Data)
|
||||
response.Type = "SetProtocolResponse"
|
||||
needSaveConfig = true
|
||||
|
||||
// 升级 Agent 命令(异步执行,不需要保存配置)
|
||||
case "UpgradeAgent":
|
||||
@@ -894,20 +882,6 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
response.Type = "UnknownCommandResponse"
|
||||
}
|
||||
|
||||
// 只有状态变更命令才保存配置
|
||||
if needSaveConfig {
|
||||
if saveErr := saveConfig(); saveErr != nil {
|
||||
fmt.Printf("❌ 保存配置失败: %v\n", saveErr)
|
||||
if err == nil {
|
||||
err = fmt.Errorf("保存配置失败: %v", saveErr)
|
||||
} else {
|
||||
err = fmt.Errorf("%v; 保存配置失败: %v", err, saveErr)
|
||||
}
|
||||
} else {
|
||||
fmt.Println("✅ 配置已保存到 gost.json")
|
||||
}
|
||||
}
|
||||
|
||||
// 发送响应
|
||||
if err != nil {
|
||||
response.Success = false
|
||||
@@ -1398,7 +1372,7 @@ func (w *WebSocketReporter) handleUpgradeAgent(data interface{}) error {
|
||||
// 执行重启脚本
|
||||
// 使用 systemd-run 在独立的 transient unit 中运行重启脚本,
|
||||
// 避免 systemctl stop 杀死 flux_agent cgroup 内所有进程(包括此脚本自身)导致 mv 未执行。
|
||||
script := fmt.Sprintf("sleep 1 && systemctl stop flux_agent && mv %s %s && systemctl start flux_agent", tmpPath, binaryPath)
|
||||
script := buildAgentRestartScript(tmpPath, binaryPath)
|
||||
cmd := exec.Command("systemd-run", "--quiet", "/bin/sh", "-c", script)
|
||||
if err := cmd.Start(); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
@@ -1432,6 +1406,14 @@ func (w *WebSocketReporter) handleRollbackAgent(data interface{}) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildAgentRestartScript(tmpPath, binaryPath string) string {
|
||||
return fmt.Sprintf(
|
||||
"sleep 1 && systemctl stop flux_agent && legacy_service='' && for service_file in /etc/systemd/system/gost.service /lib/systemd/system/gost.service /usr/lib/systemd/system/gost.service; do if [ -f \"$service_file\" ] && grep -Fq \"WorkingDirectory=/etc/gost\" \"$service_file\" && (grep -Fq \"ExecStart=/etc/gost/gost\" \"$service_file\" || (grep -Fq \"ExecStart=/usr/local/bin/gost\" \"$service_file\" && [ -f /etc/gost/config.json ] && [ -f /etc/gost/gost.json ])); then legacy_service=\"$service_file\"; break; fi; done && if [ -n \"$legacy_service\" ]; then (systemctl stop gost 2>/dev/null || true) && (systemctl disable gost 2>/dev/null || true) && rm -f /usr/local/bin/gost /etc/gost/gost \"$legacy_service\" && (systemctl daemon-reload 2>/dev/null || true); fi && mv %s %s && systemctl start flux_agent",
|
||||
tmpPath,
|
||||
binaryPath,
|
||||
)
|
||||
}
|
||||
|
||||
// updateLocalConfigJSON 将 http/tls/socks 写入工作目录下的 config.json
|
||||
func updateLocalConfigJSON(httpVal int, tlsVal int, socksVal int) error {
|
||||
path := "config.json"
|
||||
@@ -1679,17 +1661,40 @@ func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls
|
||||
candidates := buildWebSocketCandidates(addr, secret, version, http, tls, socks, "")
|
||||
fullURL := candidates[0]
|
||||
|
||||
fmt.Printf("🔗 WebSocket连接URL: %s\n", fullURL)
|
||||
fmt.Printf("🔗 WebSocket连接URL: %s\n", sanitizeWebSocketURL(fullURL))
|
||||
|
||||
reporter := NewWebSocketReporter(fullURL, secret)
|
||||
// 保存 addr, secret, version 供重连时使用
|
||||
// 保存 addr, secret, version 和协议能力供重连时使用
|
||||
reporter.addr = addr
|
||||
reporter.secret = secret
|
||||
reporter.version = version
|
||||
reporter.http = http
|
||||
reporter.tls = tls
|
||||
reporter.socks = socks
|
||||
reporter.Start()
|
||||
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)
|
||||
|
||||
@@ -1,15 +1,45 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
func captureStdout(t *testing.T, fn func()) string {
|
||||
t.Helper()
|
||||
|
||||
orig := os.Stdout
|
||||
r, w, err := os.Pipe()
|
||||
if err != nil {
|
||||
t.Fatalf("create stdout pipe: %v", err)
|
||||
}
|
||||
os.Stdout = w
|
||||
defer func() {
|
||||
os.Stdout = orig
|
||||
_ = w.Close()
|
||||
_ = r.Close()
|
||||
}()
|
||||
|
||||
fn()
|
||||
|
||||
_ = w.Close()
|
||||
|
||||
var buf bytes.Buffer
|
||||
if _, err := io.Copy(&buf, r); err != nil {
|
||||
t.Fatalf("read stdout: %v", err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func TestBuildWebSocketCandidatesSecureFirst(t *testing.T) {
|
||||
candidates := buildWebSocketCandidates("panel.example.com:443", "abc", "2.0.2", 1, 0, 1, "")
|
||||
|
||||
@@ -110,6 +140,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{
|
||||
@@ -125,3 +163,103 @@ func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
|
||||
t.Fatalf("expected response body in message, got %s", msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentUpgradeRestartScriptStopsLegacyGostService(t *testing.T) {
|
||||
script := buildAgentRestartScript("/tmp/flux_agent.new", "/etc/flux_agent/flux_agent")
|
||||
|
||||
if !strings.Contains(script, "systemctl stop flux_agent") {
|
||||
t.Fatalf("expected script to stop flux_agent, got %s", script)
|
||||
}
|
||||
if !strings.Contains(script, "mv /tmp/flux_agent.new /etc/flux_agent/flux_agent") {
|
||||
t.Fatalf("expected script to replace the flux_agent binary, got %s", script)
|
||||
}
|
||||
if !strings.Contains(script, "systemctl stop gost") {
|
||||
t.Fatalf("expected script to stop the legacy gost service, got %s", script)
|
||||
}
|
||||
if !strings.Contains(script, "systemctl disable gost") {
|
||||
t.Fatalf("expected script to disable the legacy gost service, got %s", script)
|
||||
}
|
||||
if !strings.Contains(script, "rm -f /usr/local/bin/gost") {
|
||||
t.Fatalf("expected script to remove the legacy gost binary, got %s", script)
|
||||
}
|
||||
if !strings.Contains(script, "WorkingDirectory=/etc/gost") {
|
||||
t.Fatalf("expected script to scope cleanup to the legacy FLVX gost service definition, got %s", script)
|
||||
}
|
||||
if !strings.Contains(script, "systemctl start flux_agent") {
|
||||
t.Fatalf("expected script to restart flux_agent, got %s", script)
|
||||
}
|
||||
if strings.Contains(script, "systemctl stop flux_agent && systemctl stop gost 2>/dev/null || true") {
|
||||
t.Fatalf("expected legacy gost cleanup fallback to be scoped, got %s", script)
|
||||
}
|
||||
if runtime.GOARCH == "" {
|
||||
t.Fatalf("unexpected empty runtime arch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartWebSocketReporterWithConfigPreservesProtocolDefaultsWithoutConfigFile(t *testing.T) {
|
||||
origDial := wsDial
|
||||
defer func() { wsDial = origDial }()
|
||||
|
||||
origWD, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatalf("get working directory: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = os.Chdir(origWD)
|
||||
})
|
||||
if err := os.Chdir(t.TempDir()); err != nil {
|
||||
t.Fatalf("change working directory: %v", err)
|
||||
}
|
||||
|
||||
urls := make(chan string, 1)
|
||||
wsDial = func(_ *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||
select {
|
||||
case urls <- rawURL:
|
||||
default:
|
||||
}
|
||||
return nil, nil, errors.New("dial failed")
|
||||
}
|
||||
|
||||
reporter := StartWebSocketReporterWithConfig("panel.example.com:443", "abc", 1, 0, 1, "2.0.2")
|
||||
defer reporter.Stop()
|
||||
|
||||
select {
|
||||
case rawURL := <-urls:
|
||||
if !strings.Contains(rawURL, "http=1&tls=0&socks=1") {
|
||||
t.Fatalf("expected reconnect URL to preserve startup protocol values, got %s", rawURL)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for websocket dial")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartWebSocketReporterWithConfigLogsSanitizedURL(t *testing.T) {
|
||||
origDial := wsDial
|
||||
defer func() { wsDial = origDial }()
|
||||
|
||||
ready := make(chan struct{}, 1)
|
||||
wsDial = func(_ *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||
select {
|
||||
case ready <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return nil, nil, errors.New("dial failed")
|
||||
}
|
||||
|
||||
output := captureStdout(t, func() {
|
||||
reporter := StartWebSocketReporterWithConfig("panel.example.com:443", "abc123", 1, 0, 1, "2.0.2")
|
||||
select {
|
||||
case <-ready:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("timed out waiting for websocket dial")
|
||||
}
|
||||
reporter.Stop()
|
||||
})
|
||||
|
||||
if strings.Contains(output, "secret=abc123") {
|
||||
t.Fatalf("expected logged websocket URL to mask the node secret, got %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "secret=%2A%2A%2A") {
|
||||
t.Fatalf("expected logged websocket URL to include masked secret, got %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
+79
-11
@@ -24,6 +24,11 @@ get_architecture() {
|
||||
|
||||
# 安装目录
|
||||
INSTALL_DIR="/etc/flux_agent"
|
||||
LEGACY_GOST_BINARY="/usr/local/bin/gost"
|
||||
LEGACY_GOST_CONFIG_DIR="/etc/gost"
|
||||
LEGACY_GOST_SERVICE_FILE_ETC="/etc/systemd/system/gost.service"
|
||||
LEGACY_GOST_SERVICE_FILE_LIB="/lib/systemd/system/gost.service"
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB="/usr/lib/systemd/system/gost.service"
|
||||
|
||||
# 镜像加速配置(可由面板传入或交互式询问)
|
||||
PROXY_ENABLED="${PROXY_ENABLED:-}"
|
||||
@@ -234,6 +239,69 @@ check_and_install_tcpkill() {
|
||||
return 0
|
||||
}
|
||||
|
||||
json_escape() {
|
||||
local value="$1"
|
||||
value=${value//\\/\\\\}
|
||||
value=${value//\"/\\\"}
|
||||
value=${value//$'\n'/\\n}
|
||||
value=${value//$'\r'/\\r}
|
||||
value=${value//$'\t'/\\t}
|
||||
printf '%s' "$value"
|
||||
}
|
||||
|
||||
write_flux_agent_config() {
|
||||
local path="$1"
|
||||
printf '{\n "addr": "%s",\n "secret": "%s"\n}\n' \
|
||||
"$(json_escape "$SERVER_ADDR")" \
|
||||
"$(json_escape "$SECRET")" > "$path"
|
||||
}
|
||||
|
||||
cleanup_legacy_gost_installation() {
|
||||
local matched_service_files=()
|
||||
local service_file=""
|
||||
local removed_service_file="0"
|
||||
|
||||
for service_file in "$LEGACY_GOST_SERVICE_FILE_ETC" "$LEGACY_GOST_SERVICE_FILE_LIB" "$LEGACY_GOST_SERVICE_FILE_USR_LIB"; do
|
||||
if [[ ! -f "$service_file" ]]; then
|
||||
continue
|
||||
fi
|
||||
if ! grep -Fq "WorkingDirectory=$LEGACY_GOST_CONFIG_DIR" "$service_file"; then
|
||||
continue
|
||||
fi
|
||||
if grep -Fq "ExecStart=$LEGACY_GOST_CONFIG_DIR/gost" "$service_file" || \
|
||||
(grep -Fq "ExecStart=$LEGACY_GOST_BINARY" "$service_file" && [[ -f "$LEGACY_GOST_CONFIG_DIR/config.json" && -f "$LEGACY_GOST_CONFIG_DIR/gost.json" ]]); then
|
||||
matched_service_files+=("$service_file")
|
||||
fi
|
||||
done
|
||||
|
||||
if [[ ${#matched_service_files[@]} -eq 0 ]]; then
|
||||
return 0
|
||||
fi
|
||||
|
||||
if systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
systemctl stop gost 2>/dev/null || true
|
||||
systemctl disable gost 2>/dev/null || true
|
||||
fi
|
||||
|
||||
for service_file in "${matched_service_files[@]}"; do
|
||||
if [[ -f "$service_file" ]]; then
|
||||
rm -f "$service_file"
|
||||
removed_service_file="1"
|
||||
fi
|
||||
done
|
||||
|
||||
if [[ -f "$LEGACY_GOST_BINARY" ]]; then
|
||||
rm -f "$LEGACY_GOST_BINARY"
|
||||
fi
|
||||
if [[ -f "$LEGACY_GOST_CONFIG_DIR/gost" ]]; then
|
||||
rm -f "$LEGACY_GOST_CONFIG_DIR/gost"
|
||||
fi
|
||||
|
||||
if [[ "$removed_service_file" == "1" ]]; then
|
||||
systemctl daemon-reload 2>/dev/null || true
|
||||
fi
|
||||
}
|
||||
|
||||
|
||||
# 获取用户输入的配置参数
|
||||
get_config_params() {
|
||||
@@ -279,6 +347,8 @@ install_flux_agent() {
|
||||
|
||||
mkdir -p "$INSTALL_DIR"
|
||||
|
||||
local tmp_binary="$INSTALL_DIR/flux_agent.new"
|
||||
|
||||
# 停止并禁用已有服务
|
||||
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
|
||||
echo "🔍 检测到已存在的flux_agent服务"
|
||||
@@ -286,16 +356,17 @@ install_flux_agent() {
|
||||
systemctl disable flux_agent 2>/dev/null && echo "🚫 禁用自启"
|
||||
fi
|
||||
|
||||
# 删除旧文件
|
||||
[[ -f "$INSTALL_DIR/flux_agent" ]] && echo "🧹 删除旧文件 flux_agent" && rm -f "$INSTALL_DIR/flux_agent"
|
||||
|
||||
# 下载 flux_agent
|
||||
echo "⬇️ 下载 flux_agent 中..."
|
||||
curl -L "$DOWNLOAD_URL" -o "$INSTALL_DIR/flux_agent"
|
||||
if [[ ! -f "$INSTALL_DIR/flux_agent" || ! -s "$INSTALL_DIR/flux_agent" ]]; then
|
||||
rm -f "$tmp_binary"
|
||||
curl -L "$DOWNLOAD_URL" -o "$tmp_binary"
|
||||
if [[ ! -f "$tmp_binary" || ! -s "$tmp_binary" ]]; then
|
||||
rm -f "$tmp_binary"
|
||||
echo "❌ 下载失败,请检查网络或下载链接。"
|
||||
exit 1
|
||||
fi
|
||||
cleanup_legacy_gost_installation
|
||||
mv "$tmp_binary" "$INSTALL_DIR/flux_agent"
|
||||
chmod +x "$INSTALL_DIR/flux_agent"
|
||||
echo "✅ 下载完成"
|
||||
|
||||
@@ -305,12 +376,7 @@ install_flux_agent() {
|
||||
# 写入 config.json (安装时总是创建新的)
|
||||
CONFIG_FILE="$INSTALL_DIR/config.json"
|
||||
echo "📄 创建新配置: config.json"
|
||||
cat > "$CONFIG_FILE" <<EOF
|
||||
{
|
||||
"addr": "$SERVER_ADDR",
|
||||
"secret": "$SECRET"
|
||||
}
|
||||
EOF
|
||||
write_flux_agent_config "$CONFIG_FILE"
|
||||
|
||||
# 写入 gost.json
|
||||
GOST_CONFIG="$INSTALL_DIR/gost.json"
|
||||
@@ -380,11 +446,13 @@ update_flux_agent() {
|
||||
|
||||
# 先下载新版本
|
||||
echo "⬇️ 下载最新版本..."
|
||||
rm -f "$INSTALL_DIR/flux_agent.new"
|
||||
curl -L "$DOWNLOAD_URL" -o "$INSTALL_DIR/flux_agent.new"
|
||||
if [[ ! -f "$INSTALL_DIR/flux_agent.new" || ! -s "$INSTALL_DIR/flux_agent.new" ]]; then
|
||||
echo "❌ 下载失败。"
|
||||
return 1
|
||||
fi
|
||||
cleanup_legacy_gost_installation
|
||||
|
||||
# 停止服务
|
||||
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user