mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 23:56:36 +08:00
Compare commits
26 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 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 |
@@ -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,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`.
|
||||
@@ -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
|
||||
}
|
||||
@@ -1659,7 +1693,7 @@ func compactErrorMessage(msg string) string {
|
||||
return strings.Join(strings.Fields(strings.ToLower(msg)), "")
|
||||
}
|
||||
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, 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,14 +1736,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 {
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
handlerConfig := service["handler"].(map[string]interface{})
|
||||
if handlerConfig["metadata"] == nil {
|
||||
handlerConfig["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
service["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
|
||||
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
|
||||
}
|
||||
if protocol == "udp" {
|
||||
listenerMetadata := map[string]interface{}{
|
||||
@@ -1727,9 +1765,6 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
}
|
||||
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
|
||||
}
|
||||
if limiterID != nil && *limiterID > 0 {
|
||||
service["limiter"] = strconv.FormatInt(*limiterID, 10)
|
||||
}
|
||||
services = append(services, service)
|
||||
}
|
||||
|
||||
@@ -1829,22 +1864,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)
|
||||
}
|
||||
@@ -1852,15 +1881,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)},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1888,3 +1961,20 @@ func (h *Handler) upsertLimiterOnNode(nodeID int64, limiterID int64, speed int)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) ensureTrafficLimiterOnNode(nodeID int64, name string, totalSpeed *int, ipSpeed *int) error {
|
||||
payload := buildTrafficLimiterPayload(name, totalSpeed, ipSpeed)
|
||||
limits, _ := payload["limits"].([]string)
|
||||
if name == "" || len(limits) == 0 {
|
||||
return nil
|
||||
}
|
||||
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
|
||||
if !isAlreadyExistsMessage(err.Error()) {
|
||||
return fmt.Errorf("限速规则下发失败: %w", err)
|
||||
}
|
||||
if _, updateErr := h.sendNodeCommand(nodeID, "UpdateLimiters", buildLimiterUpdatePayload(name, payload), false, false); updateErr != nil {
|
||||
return fmt.Errorf("限速规则更新失败: %w", updateErr)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -378,7 +378,7 @@ func TestRetryTunnelServiceAddWithCleanupReturnsCleanupError(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22000, "10.9.8.7", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -393,7 +393,7 @@ func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -409,7 +409,7 @@ func TestBuildForwardServiceConfigs_DefaultListenAddrWhenBindIPEmpty(t *testing.
|
||||
func TestBuildForwardServiceConfigs_BindIPAlreadyContainsPort(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 55555, "3.3.3.3:12345", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -464,7 +464,7 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, nil, "")
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, tt.port, tt.bindIP, forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
@@ -478,6 +478,58 @@ func TestBuildForwardServiceConfigs_IPv6BindIP(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildConnLimiterConfigCombinesTotalAndPerIP(t *testing.T) {
|
||||
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, MaxConn: 100, IPMaxConn: 5}, 37)
|
||||
want := []forwardLimiterConfig{{Name: "rule_conn_limit_42", Limits: []string{"$ 100", "$$ 5"}}}
|
||||
if !reflect.DeepEqual(cfgs, want) {
|
||||
t.Fatalf("expected %+v, got %+v", want, cfgs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildConnLimiterConfigUsesUserTotalWithRulePerIP(t *testing.T) {
|
||||
cfgs := buildConnLimiterConfigs(&forwardRecord{ID: 42, UserID: 9, IPMaxConn: 5}, 37)
|
||||
want := []forwardLimiterConfig{
|
||||
{Name: "user_conn_limit_9", Limits: []string{"$ 37"}},
|
||||
{Name: "rule_conn_limit_42", Limits: []string{"$$ 5"}},
|
||||
}
|
||||
if !reflect.DeepEqual(cfgs, want) {
|
||||
t.Fatalf("expected %+v, got %+v", want, cfgs)
|
||||
}
|
||||
if got := joinLimiterNames(cfgs); got != "user_conn_limit_9,rule_conn_limit_42" {
|
||||
t.Fatalf("expected composite limiter names, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildTrafficLimiterPayloadUsesOnlyPerIPRulesWhenTotalIsSeparate(t *testing.T) {
|
||||
payload := buildTrafficLimiterPayload("rule_traffic_limit_42", nil, intPtr(40))
|
||||
wantLimits := []string{"0.0.0.0/0 5.0MB 5.0MB", "::/0 5.0MB 5.0MB"}
|
||||
if payload["name"] != "rule_traffic_limit_42" {
|
||||
t.Fatalf("expected name rule_traffic_limit_42, got %v", payload["name"])
|
||||
}
|
||||
if !reflect.DeepEqual(payload["limits"], wantLimits) {
|
||||
t.Fatalf("expected limits %v, got %v", wantLimits, payload["limits"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigsUsesRuntimeLimiterNames(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "0.0.0.0", UDPListenAddr: "[::]"}
|
||||
services := buildForwardServiceConfigs("1_2_0", forward, nil, node, 22001, "", forwardRuntimeLimiters{TrafficLimiter: "rule_traffic_limit_42", ConnLimiter: "rule_conn_limit_42"})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
for _, service := range services {
|
||||
if service["limiter"] != "rule_traffic_limit_42" {
|
||||
t.Fatalf("expected traffic limiter rule_traffic_limit_42, got %v", service["limiter"])
|
||||
}
|
||||
if service["climiter"] != "rule_conn_limit_42" {
|
||||
t.Fatalf("expected conn limiter rule_conn_limit_42, got %v", service["climiter"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func intPtr(v int) *int { return &v }
|
||||
|
||||
func TestProcessServerAddress_StripsURLSchemeAndPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -8,7 +9,7 @@ import (
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildForwardServiceConfigsPreservesProxyProtocolWithInterfaceMetadata(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
@@ -24,21 +25,33 @@ func TestBuildForwardServiceConfigsPreservesProxyProtocolWithInterfaceMetadata(t
|
||||
UDPListenAddr: "0.0.0.0",
|
||||
}
|
||||
|
||||
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", nil, "")
|
||||
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 {
|
||||
metadata, ok := service["metadata"].(map[string]interface{})
|
||||
serviceMetadata, ok := service["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected metadata map, got %T", service["metadata"])
|
||||
}
|
||||
if metadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface metadata eth0, got %v", metadata["interface"])
|
||||
if serviceMetadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface metadata eth0, got %v", serviceMetadata["interface"])
|
||||
}
|
||||
if metadata["proxyProtocol"] != 2 {
|
||||
t.Fatalf("expected proxyProtocol 2, got %v", metadata["proxyProtocol"])
|
||||
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"])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -61,6 +74,8 @@ func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
|
||||
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)
|
||||
@@ -69,6 +84,8 @@ func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
|
||||
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 {
|
||||
@@ -85,14 +102,22 @@ func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
|
||||
RemoteAddr: "9.9.9.9:443",
|
||||
Strategy: "fifo",
|
||||
Status: 1,
|
||||
IPMaxConn: 5,
|
||||
IPSpeedID: sql.NullInt64{Int64: 21, Valid: true},
|
||||
ProxyProtocol: 2,
|
||||
}, nil)
|
||||
|
||||
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: %v", err)
|
||||
var record model.Forward
|
||||
if err := r.DB().Where("id = ?", forwardID).First(&record).Error; err != nil {
|
||||
t.Fatalf("query forward: %v", err)
|
||||
}
|
||||
if proxyProtocol != 2 {
|
||||
t.Fatalf("expected proxyProtocol restored to 2, got %d", proxyProtocol)
|
||||
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,17 @@ type Handler struct {
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
|
||||
upgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
upgradeMu sync.Mutex
|
||||
pendingUpgradeRedeploy map[int64]struct{}
|
||||
nodeOnlineRedeployAt map[int64]time.Time
|
||||
nodeOnlineRedeployQueued map[int64]struct{}
|
||||
nodeOnlineRedeploying map[int64]struct{}
|
||||
|
||||
qualityProber *tunnelQualityProber
|
||||
}
|
||||
|
||||
const monitorTunnelQualityEnabledConfigKey = "monitor_tunnel_quality_enabled"
|
||||
const allowLocalRemoteAddrConfigKey = "allow_local_remote_addr"
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
@@ -95,13 +101,16 @@ const (
|
||||
|
||||
func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
h := &Handler{
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
repo: repo,
|
||||
jwtSecret: jwtSecret,
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
nodeOnlineRedeployAt: make(map[int64]time.Time),
|
||||
nodeOnlineRedeployQueued: make(map[int64]struct{}),
|
||||
nodeOnlineRedeploying: make(map[int64]struct{}),
|
||||
}
|
||||
h.healthCheck = health.NewChecker(repo, h.wsServer)
|
||||
h.qualityProber = newTunnelQualityProber(h)
|
||||
@@ -144,6 +153,7 @@ 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/license/activate", h.licenseActivate)
|
||||
mux.HandleFunc("/api/v1/backup/export", h.backupExport)
|
||||
mux.HandleFunc("/api/v1/backup/import", h.backupImport)
|
||||
@@ -796,11 +806,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 +1038,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 +1058,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("请求失败"))
|
||||
|
||||
@@ -138,6 +138,15 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("请不要作死"))
|
||||
return
|
||||
}
|
||||
oldUser, err := h.repo.GetUserByID(id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if oldUser == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("用户不存在"))
|
||||
return
|
||||
}
|
||||
|
||||
dup, err := h.repo.UserExistsExcluding(username, id)
|
||||
if err != nil {
|
||||
@@ -210,6 +219,17 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
}
|
||||
if oldUser.MaxConn != maxConn {
|
||||
warnings, syncErr := h.syncUserMaxConnForwards(id)
|
||||
if syncErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(fmt.Sprintf("最大连接数下发失败: %v", syncErr)))
|
||||
return
|
||||
}
|
||||
if len(warnings) > 0 {
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"warnings": warnings}))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
@@ -1726,11 +1746,13 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
|
||||
return
|
||||
}
|
||||
if roleID != 0 {
|
||||
if roleID != 0 && !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
if roleID != 0 {
|
||||
if speedIDVal, ok := req["speedId"]; ok && speedIDVal != nil {
|
||||
response.WriteJSON(w, response.Err(-1, "普通用户无法设置限速规则"))
|
||||
return
|
||||
@@ -1742,6 +1764,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if roleID != 0 {
|
||||
if ipSpeedIDVal, ok := req["ipSpeedId"]; ok && ipSpeedIDVal != nil {
|
||||
response.WriteJSON(w, response.Err(-1, "普通用户无法设置每 IP 限速规则"))
|
||||
return
|
||||
}
|
||||
}
|
||||
ipSpeedID := asAnyToInt64Ptr(req["ipSpeedId"])
|
||||
ipSpeedID, err = h.normalizeSpeedLimitReference(ipSpeedID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
port = h.pickTunnelPort(tunnelID)
|
||||
@@ -1780,9 +1814,13 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
userName = "user"
|
||||
}
|
||||
maxConn := asInt(req["maxConn"], 0)
|
||||
ipMaxConn := asInt(req["ipMaxConn"], 0)
|
||||
if ipMaxConn < 0 {
|
||||
ipMaxConn = 0
|
||||
}
|
||||
proxyProtocol := asInt(req["proxyProtocol"], 0)
|
||||
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, proxyProtocol)
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, ipMaxConn, nullableInt(ipSpeedID), proxyProtocol)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1853,7 +1891,7 @@ 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, err.Error()))
|
||||
return
|
||||
@@ -1881,6 +1919,27 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
} else if _, ok := req["speedId"]; ok {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
rawIPSpeedID, hasIPSpeedID := req["ipSpeedId"]
|
||||
requestedIPSpeedID := asAnyToInt64Ptr(rawIPSpeedID)
|
||||
newIPSpeedID := forward.IPSpeedID
|
||||
if actorRole != 0 {
|
||||
if hasIPSpeedID && !sameSpeedLimitSelection(forward.IPSpeedID, requestedIPSpeedID) {
|
||||
response.WriteJSON(w, response.Err(-1, "普通用户无法修改每 IP 限速规则"))
|
||||
return
|
||||
}
|
||||
} else {
|
||||
ipSpeedID := requestedIPSpeedID
|
||||
ipSpeedID, err = h.normalizeSpeedLimitReference(ipSpeedID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if ipSpeedID != nil {
|
||||
newIPSpeedID = sql.NullInt64{Int64: *ipSpeedID, Valid: true}
|
||||
} else if hasIPSpeedID {
|
||||
newIPSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
}
|
||||
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
@@ -1934,9 +1993,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, proxyProtocol); 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
|
||||
}
|
||||
@@ -4088,7 +4151,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
||||
h.repo.RollbackForwardFields(
|
||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.ProxyProtocol,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
@@ -4252,6 +4315,26 @@ func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncUserMaxConnForwards(userID int64) ([]string, error) {
|
||||
forwards, err := h.listActiveForwardsByUser(userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
warnings := make([]string, 0)
|
||||
for i := range forwards {
|
||||
f := &forwards[i]
|
||||
if f.MaxConn > 0 {
|
||||
continue
|
||||
}
|
||||
syncWarnings, syncErr := h.syncForwardServicesWithWarnings(f, "UpdateService", true)
|
||||
warnings = append(warnings, syncWarnings...)
|
||||
if syncErr != nil {
|
||||
return warnings, syncErr
|
||||
}
|
||||
}
|
||||
return warnings, nil
|
||||
}
|
||||
|
||||
// cleanupForwardsForUserTunnel deletes all forwarding rules belonging to a
|
||||
// specific user+tunnel pair. It notifies nodes to remove the runtime services
|
||||
// first, then deletes the DB records. This is best-effort: individual failures
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -128,12 +128,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 +152,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)
|
||||
}
|
||||
@@ -287,7 +294,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 {
|
||||
|
||||
@@ -39,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"`
|
||||
@@ -396,23 +398,120 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
|
||||
}
|
||||
|
||||
func (h *Handler) onNodeOnline(nodeID int64) {
|
||||
h.consumeNodePendingUpgradeRedeploy(nodeID)
|
||||
// Always redeploy rules on reconnection, not just for pending upgrade nodes.
|
||||
// This handles cases where the node restarted and lost its in-memory config
|
||||
// before persistence had time to flush, or if the panel also restarted.
|
||||
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
|
||||
@@ -442,7 +541,7 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
}
|
||||
|
||||
// Retry failed items with exponential backoff (max 3 attempts)
|
||||
h.retryFailedRedeploys(nodeID, tunnelFailed, failedForwards)
|
||||
return h.retryFailedRedeploys(nodeID, tunnelFailed, failedForwards)
|
||||
}
|
||||
|
||||
// isRetryableError returns true if the error looks transient and worth retrying.
|
||||
@@ -463,9 +562,9 @@ func isRetryableError(err error) bool {
|
||||
}
|
||||
|
||||
// retryFailedRedeploys retries failed tunnels and forwards with exponential backoff.
|
||||
func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]struct{}, failedForwards []failedForward) {
|
||||
func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]struct{}, failedForwards []failedForward) bool {
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
return
|
||||
return true
|
||||
}
|
||||
|
||||
const maxRetries = 3
|
||||
@@ -507,7 +606,7 @@ func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]stru
|
||||
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
fmt.Printf("post-upgrade redeploy retry: all items recovered on node %d\n", nodeID)
|
||||
return
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -518,4 +617,5 @@ func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]stru
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -45,6 +45,8 @@ type Forward struct {
|
||||
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"`
|
||||
}
|
||||
|
||||
@@ -428,20 +430,22 @@ 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"`
|
||||
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"`
|
||||
}
|
||||
@@ -532,16 +536,18 @@ 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
|
||||
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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -690,11 +790,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
|
||||
@@ -725,7 +825,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,
|
||||
"maxConn": u.MaxConn,
|
||||
}
|
||||
if quota := quotaMap[u.ID]; quota != nil {
|
||||
item["dailyQuotaGB"] = quota.DailyLimitGB
|
||||
@@ -766,29 +866,33 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
}
|
||||
|
||||
type fwdRow struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
TunnelName string
|
||||
TrafficRatio float64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
CreatedTime int64
|
||||
Status int
|
||||
Inx int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
ProxyProtocol int
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
TunnelName string
|
||||
TrafficRatio float64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
CreatedTime int64
|
||||
Status int
|
||||
Inx int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
IPSpeedLimitName string
|
||||
ProxyProtocol int
|
||||
}
|
||||
|
||||
var rows []fwdRow
|
||||
err := r.db.Model(&model.Forward{}).
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.proxy_protocol").
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.ip_max_conn, forward.ip_speed_id, COALESCE(ip_speed_limit.name, '') AS ip_speed_limit_name, forward.proxy_protocol").
|
||||
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 {
|
||||
@@ -809,12 +913,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,
|
||||
"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
|
||||
@@ -2003,8 +2114,17 @@ 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 {
|
||||
return nil, err
|
||||
@@ -2384,30 +2504,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", "proxy_protocol",
|
||||
"in_flow", "out_flow", "updated_time", "status", "inx", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
if err != nil {
|
||||
@@ -3355,7 +3485,7 @@ func (r *Repository) GetNodeMetrics(nodeID int64, startMs, endMs int64) ([]model
|
||||
|
||||
rangeMs := endMs - startMs
|
||||
const maxRawRangeMs = int64(60 * 60 * 1000) // 1 hour — return raw data for short ranges
|
||||
const targetPoints = 500 // target number of chart points for downsampled data
|
||||
const targetPoints = 500 // target number of chart points for downsampled data
|
||||
|
||||
// For short ranges, return raw data (full resolution).
|
||||
if rangeMs <= maxRawRangeMs {
|
||||
|
||||
@@ -45,15 +45,19 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
rows = append(rows, model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -64,7 +68,6 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
|
||||
func (r *Repository) ListActiveTunnelIDsByNode(nodeID int64) ([]int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
@@ -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 {
|
||||
@@ -134,6 +230,8 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
}
|
||||
if strings.TrimSpace(fr.Strategy) == "" {
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestChunkFlowUploadForwardIDs(t *testing.T) {
|
||||
ids := make([]int64, 0, 1001)
|
||||
for i := int64(1); i <= 1001; i++ {
|
||||
ids = append(ids, i)
|
||||
}
|
||||
|
||||
chunks := chunkFlowUploadForwardIDs(ids)
|
||||
if len(chunks) != 3 {
|
||||
t.Fatalf("expected 3 chunks, got %d", len(chunks))
|
||||
}
|
||||
if len(chunks[0]) != 500 || len(chunks[1]) != 500 || len(chunks[2]) != 1 {
|
||||
t.Fatalf("unexpected chunk sizes: %d, %d, %d", len(chunks[0]), len(chunks[1]), len(chunks[2]))
|
||||
}
|
||||
if chunks[0][0] != 1 || chunks[1][0] != 501 || chunks[2][0] != 1001 {
|
||||
t.Fatalf("unexpected chunk boundaries: %#v %#v %#v", chunks[0][:1], chunks[1][:1], chunks[2][:1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSortedFlowUploadTargetIDs(t *testing.T) {
|
||||
totals := map[int64][2]int64{
|
||||
9: {1, 1},
|
||||
2: {1, 1},
|
||||
7: {1, 1},
|
||||
}
|
||||
|
||||
got := sortedFlowUploadTargetIDs(totals)
|
||||
want := []int64{2, 7, 9}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("expected sorted ids %v, got %v", want, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
metas, err := r.GetFlowUploadForwardMetas([]int64{20, 99})
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
t.Fatalf("unexpected meta for forward 20: %#v", metas[20])
|
||||
}
|
||||
if _, ok := metas[99]; ok {
|
||||
t.Fatalf("did not expect meta for missing forward 99")
|
||||
}
|
||||
|
||||
err = r.ApplyFlowUploadDeltasBatch([]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}})
|
||||
if err != nil {
|
||||
t.Fatalf("apply flow batch: %v", err)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
|
||||
t.Fatalf("expected forward in_flow=480, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
|
||||
t.Fatalf("expected user out_flow=660, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
|
||||
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Exec(`INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx) VALUES(25, 2, 'u2', 'f25', 99, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
metas, err := r.GetFlowUploadForwardMetas([]int64{25})
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
meta, ok := metas[25]
|
||||
if !ok {
|
||||
t.Fatalf("expected metadata for forward with missing tunnel row")
|
||||
}
|
||||
if meta.ForwardID != 25 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
|
||||
t.Fatalf("unexpected fallback meta: %#v", meta)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddUserQuotaUsageBatchReturnsNormalizedViews(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "quota-batch.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
if err := r.DB().Exec(`INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status) VALUES(2, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
views, err := r.AddUserQuotaUsageBatch(map[int64]int64{2: 1140}, now)
|
||||
if err != nil {
|
||||
t.Fatalf("batch quota update: %v", err)
|
||||
}
|
||||
if views[2] == nil || views[2].DailyUsedBytes != 1140 || views[2].MonthlyUsedBytes != 1140 {
|
||||
t.Fatalf("unexpected quota view: %#v", views[2])
|
||||
}
|
||||
}
|
||||
|
||||
func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...interface{}) int64 {
|
||||
t.Helper()
|
||||
var value int64
|
||||
if err := r.DB().Raw(query, args...).Row().Scan(&value); err != nil {
|
||||
t.Fatalf("query %q failed: %v", query, err)
|
||||
}
|
||||
return value
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -46,6 +47,205 @@ func TestGetForwardRecordIncludesProxyProtocol(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
|
||||
@@ -695,7 +695,7 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
|
||||
return p
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, proxyProtocol int) error {
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -708,6 +708,8 @@ func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remote
|
||||
"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
|
||||
@@ -783,24 +785,26 @@ func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int,
|
||||
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, proxyProtocol int, now int64) {
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int, 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,
|
||||
"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,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -1260,7 +1264,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
return ut.ID, true, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, proxyProtocol int) (int64, error) {
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, ipMaxConn int, ipSpeedID interface{}, proxyProtocol int) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
@@ -1281,6 +1285,8 @@ func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnel
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
IPMaxConn: ipMaxConn,
|
||||
IPSpeedID: nullInt64FromInterface(ipSpeedID),
|
||||
ProxyProtocol: proxyProtocol,
|
||||
}
|
||||
if err := tx.Create(&fwd).Error; err != nil {
|
||||
|
||||
@@ -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,12 @@ 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)
|
||||
@@ -194,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)
|
||||
@@ -217,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()
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -234,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",
|
||||
})
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -11,6 +11,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 {
|
||||
@@ -209,6 +210,7 @@ func deleteServices(req deleteServicesRequest) error {
|
||||
}
|
||||
return nil
|
||||
})
|
||||
xservice.GetGlobalTrafficManager().RemoveServices(namesToRemove...)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -145,11 +145,12 @@ type ServiceMonitorCheckResult struct {
|
||||
}
|
||||
|
||||
const (
|
||||
reporterReadWait = 60 * time.Second
|
||||
reporterWriteWait = 5 * time.Second
|
||||
wsPingInterval = 20 * time.Second // 独立 WebSocket ping 间隔
|
||||
initialBackoff = 2 * time.Second // 重连初始退避
|
||||
maxBackoff = 2 * time.Minute // 重连最大退避
|
||||
reporterReadWait = 60 * time.Second
|
||||
reporterWriteWait = 5 * time.Second
|
||||
wsPingInterval = 20 * time.Second // 独立 WebSocket ping 间隔
|
||||
initialBackoff = 2 * time.Second // 重连初始退避
|
||||
maxBackoff = 2 * time.Minute // 重连最大退避
|
||||
defaultMetricReportInterval = 5 * time.Second
|
||||
)
|
||||
|
||||
type WebSocketReporter struct {
|
||||
@@ -189,9 +190,9 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
|
||||
|
||||
return &WebSocketReporter{
|
||||
url: serverURL,
|
||||
curBackoff: initialBackoff, // 当前退避间隔
|
||||
pingInterval: 1 * time.Second, // 指标上报间隔(每秒采集)
|
||||
configInterval: 10 * time.Minute, // 配置上报间隔
|
||||
curBackoff: initialBackoff, // 当前退避间隔
|
||||
pingInterval: defaultMetricReportInterval, // 指标上报间隔
|
||||
configInterval: 10 * time.Minute, // 配置上报间隔
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
connected: false,
|
||||
|
||||
@@ -110,6 +110,14 @@ func TestSanitizeWebSocketURL(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewWebSocketReporterUsesReducedMetricInterval(t *testing.T) {
|
||||
reporter := NewWebSocketReporter("panel.example.com:443", "abc")
|
||||
|
||||
if reporter.pingInterval != defaultMetricReportInterval {
|
||||
t.Fatalf("expected metric interval %s, got %s", defaultMetricReportInterval, reporter.pingInterval)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
|
||||
err := errors.New("websocket: bad handshake")
|
||||
resp := &http.Response{
|
||||
|
||||
@@ -41,6 +41,7 @@ import type {
|
||||
MonitorPermissionApiItem,
|
||||
MonitorAccessApiData,
|
||||
TunnelQualityApiItem,
|
||||
StorageSummaryApiData,
|
||||
} from "./types";
|
||||
|
||||
import axios from "axios";
|
||||
@@ -254,6 +255,9 @@ export const updateConfigs = (configMap: Record<string, string>) =>
|
||||
export const updateConfig = (name: string, value: string) =>
|
||||
Network.post("/config/update-single", { name, value });
|
||||
|
||||
export const getStorageSummary = () =>
|
||||
Network.get<StorageSummaryApiData>("/system/storage");
|
||||
|
||||
export const activateLicense = (licenseKey: string) =>
|
||||
Network.post("/license/activate", { license_key: licenseKey });
|
||||
|
||||
@@ -418,8 +422,11 @@ export interface AnnouncementData {
|
||||
|
||||
export const getAnnouncement = () =>
|
||||
Network.get<AnnouncementData>("/announcement/get");
|
||||
export const updateAnnouncement = (data: AnnouncementData) =>
|
||||
Network.post("/announcement/update", data);
|
||||
export const updateAnnouncement = ({
|
||||
content,
|
||||
enabled,
|
||||
}: Pick<AnnouncementData, "content" | "enabled">) =>
|
||||
Network.post("/announcement/update", { content, enabled });
|
||||
|
||||
export const getNodeMetrics = (
|
||||
nodeId: number,
|
||||
|
||||
@@ -71,6 +71,9 @@ export interface ForwardApiItem {
|
||||
userId?: number;
|
||||
tunnelId?: number;
|
||||
speedId?: number | null;
|
||||
ipMaxConn?: number;
|
||||
ipSpeedId?: number | null;
|
||||
ipSpeedLimitName?: string;
|
||||
maxConn?: number;
|
||||
proxyProtocol?: number;
|
||||
inx?: number;
|
||||
@@ -372,6 +375,8 @@ export interface ForwardMutationPayload {
|
||||
remoteAddr?: string;
|
||||
strategy?: string;
|
||||
speedId?: number | null;
|
||||
ipMaxConn?: number;
|
||||
ipSpeedId?: number | null;
|
||||
maxConn?: number;
|
||||
proxyProtocol?: number;
|
||||
}
|
||||
@@ -469,6 +474,12 @@ export interface ServiceMonitorLimitsApiData {
|
||||
maxTimeoutSec: number;
|
||||
}
|
||||
|
||||
export interface StorageSummaryApiData {
|
||||
dbType: string;
|
||||
databaseSizeBytes: number;
|
||||
databaseSizeText: string;
|
||||
}
|
||||
|
||||
export interface MonitorNodeApiItem {
|
||||
id: number;
|
||||
inx: number;
|
||||
|
||||
@@ -26,6 +26,7 @@ import {
|
||||
importBackup,
|
||||
getAnnouncement,
|
||||
updateAnnouncement,
|
||||
getStorageSummary,
|
||||
type AnnouncementData,
|
||||
} from "@/api";
|
||||
import { BackIcon, SettingsIcon } from "@/components/icons";
|
||||
@@ -146,6 +147,14 @@ const CONFIG_ITEMS: ConfigItem[] = [
|
||||
"关闭后,前端停止自动刷新,后端停止实时隧道质量探测(全局配置)",
|
||||
type: "switch",
|
||||
},
|
||||
{
|
||||
key: "monitor_retention_days",
|
||||
label: "监控数据保留天数",
|
||||
placeholder: "7",
|
||||
description:
|
||||
"统一清理节点指标、隧道流量、服务监控结果和隧道质量历史;默认 7 天。",
|
||||
type: "input",
|
||||
},
|
||||
{
|
||||
key: "captcha_enabled",
|
||||
label: "启用验证码",
|
||||
@@ -185,6 +194,13 @@ const CONFIG_ITEMS: ConfigItem[] = [
|
||||
dependsOn: "github_proxy_enabled",
|
||||
dependsValue: "true",
|
||||
},
|
||||
{
|
||||
key: "allow_local_remote_addr",
|
||||
label: "允许转发到本地地址",
|
||||
description:
|
||||
"开启后,普通用户创建或编辑规则时可将目标地址指向 127.0.0.1、10.x.x.x、172.16-31.x.x、192.168.x.x 等本地或内网地址。默认关闭以降低开放代理风险。",
|
||||
type: "switch",
|
||||
},
|
||||
];
|
||||
|
||||
const BACKUP_TYPE_OPTIONS = [
|
||||
@@ -213,12 +229,14 @@ const getInitialConfigs = (): Record<string, string> => {
|
||||
"cloudflare_secret_key",
|
||||
"forward_compact_mode",
|
||||
"monitor_tunnel_quality_enabled",
|
||||
"monitor_retention_days",
|
||||
"ip",
|
||||
"panel_domain",
|
||||
"app_logo",
|
||||
"app_favicon",
|
||||
"github_proxy_enabled",
|
||||
"github_proxy_url",
|
||||
"allow_local_remote_addr",
|
||||
];
|
||||
const initialConfigs: Record<string, string> = {};
|
||||
|
||||
@@ -280,6 +298,7 @@ export default function ConfigPage() {
|
||||
const [brandUploading, setBrandUploading] = useState<
|
||||
Partial<Record<BrandPreviewKey, boolean>>
|
||||
>({});
|
||||
const [storageSummary, setStorageSummary] = useState("加载中...");
|
||||
|
||||
const canGoBack =
|
||||
typeof window !== "undefined" &&
|
||||
@@ -339,10 +358,26 @@ export default function ConfigPage() {
|
||||
}
|
||||
};
|
||||
|
||||
const loadStorageSummary = async () => {
|
||||
try {
|
||||
const response = await getStorageSummary();
|
||||
|
||||
if (response.code === 0 && response.data?.databaseSizeText) {
|
||||
setStorageSummary(response.data.databaseSizeText);
|
||||
|
||||
return;
|
||||
}
|
||||
setStorageSummary("获取失败");
|
||||
} catch {
|
||||
setStorageSummary("获取失败");
|
||||
}
|
||||
};
|
||||
|
||||
useEffect(() => {
|
||||
const timer = setTimeout(() => {
|
||||
loadConfigs(initialConfigs);
|
||||
loadAnnouncement();
|
||||
loadStorageSummary();
|
||||
}, 100);
|
||||
|
||||
return () => clearTimeout(timer);
|
||||
@@ -1316,6 +1351,22 @@ export default function ConfigPage() {
|
||||
</Select>
|
||||
</div>
|
||||
|
||||
<Divider className="my-2" />
|
||||
|
||||
<div className="space-y-3">
|
||||
<div className="flex flex-col gap-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>
|
||||
<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>
|
||||
|
||||
<div className="flex justify-end pt-6 border-t border-divider/50 mt-4">
|
||||
<Button
|
||||
color="primary"
|
||||
|
||||
@@ -360,11 +360,7 @@ export const useDashboardData = (): DashboardDataState => {
|
||||
if (updateTime > storedTime) {
|
||||
setIsAnnouncementModalOpen(true);
|
||||
}
|
||||
} catch (err) {
|
||||
console.warn(
|
||||
"Failed to read localStorage for announcement state",
|
||||
err,
|
||||
);
|
||||
} catch {
|
||||
setIsAnnouncementModalOpen(true);
|
||||
}
|
||||
} else {
|
||||
@@ -385,8 +381,8 @@ export const useDashboardData = (): DashboardDataState => {
|
||||
"flvx_announcement_seen_time",
|
||||
announcement.update_time.toString(),
|
||||
);
|
||||
} catch (err) {
|
||||
console.warn("Failed to set localStorage for announcement state", err);
|
||||
} catch {
|
||||
// Ignore localStorage write failures and keep the modal dismissed.
|
||||
}
|
||||
}
|
||||
}, [announcement]);
|
||||
|
||||
@@ -125,6 +125,9 @@ interface Forward {
|
||||
userId?: number;
|
||||
inx?: number;
|
||||
speedId?: number | null;
|
||||
ipMaxConn?: number;
|
||||
ipSpeedId?: number | null;
|
||||
ipSpeedLimitName?: string;
|
||||
proxyProtocol?: number;
|
||||
}
|
||||
|
||||
@@ -160,6 +163,8 @@ interface ForwardForm {
|
||||
interfaceName?: string;
|
||||
strategy: string;
|
||||
speedId: number | null;
|
||||
ipMaxConn?: number;
|
||||
ipSpeedId: number | null;
|
||||
maxConn?: number;
|
||||
proxyProtocol?: number;
|
||||
}
|
||||
@@ -578,6 +583,16 @@ const mapForwardApiItems = (items: ForwardApiItem[]): Forward[] => {
|
||||
typeof forward.speedId === "number" || forward.speedId === null
|
||||
? forward.speedId
|
||||
: undefined,
|
||||
ipMaxConn:
|
||||
typeof forward.ipMaxConn === "number" ? forward.ipMaxConn : undefined,
|
||||
ipSpeedId:
|
||||
typeof forward.ipSpeedId === "number" || forward.ipSpeedId === null
|
||||
? forward.ipSpeedId
|
||||
: undefined,
|
||||
ipSpeedLimitName:
|
||||
typeof forward.ipSpeedLimitName === "string"
|
||||
? forward.ipSpeedLimitName
|
||||
: undefined,
|
||||
maxConn: typeof forward.maxConn === "number" ? forward.maxConn : undefined,
|
||||
proxyProtocol:
|
||||
typeof forward.proxyProtocol === "number"
|
||||
@@ -1316,6 +1331,8 @@ export default function ForwardPage() {
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
ipMaxConn: 0,
|
||||
ipSpeedId: null,
|
||||
maxConn: 0,
|
||||
proxyProtocol: 0,
|
||||
});
|
||||
@@ -2030,6 +2047,7 @@ export default function ForwardPage() {
|
||||
};
|
||||
|
||||
const selectedSpeedId = normalizeSpeedId(form.speedId);
|
||||
const selectedIPSpeedId = normalizeSpeedId(form.ipSpeedId);
|
||||
|
||||
const validateForm = (): boolean => {
|
||||
const newErrors: { [key: string]: string } = {};
|
||||
@@ -2105,6 +2123,8 @@ export default function ForwardPage() {
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
ipMaxConn: 0,
|
||||
ipSpeedId: null,
|
||||
proxyProtocol: 0,
|
||||
});
|
||||
setErrors({});
|
||||
@@ -2126,6 +2146,8 @@ export default function ForwardPage() {
|
||||
interfaceName: forward.interfaceName || "",
|
||||
strategy: forward.strategy || "fifo",
|
||||
speedId: normalizeSpeedId(forward.speedId),
|
||||
ipMaxConn: forward.ipMaxConn ?? 0,
|
||||
ipSpeedId: normalizeSpeedId(forward.ipSpeedId),
|
||||
maxConn: forward.maxConn ?? 0,
|
||||
proxyProtocol: forward.proxyProtocol ?? 0,
|
||||
});
|
||||
@@ -2245,6 +2267,8 @@ export default function ForwardPage() {
|
||||
let res: { code: number; msg: string };
|
||||
const normalizedSpeedId = normalizeSpeedId(form.speedId);
|
||||
const speedLimitAutoCleared = isMissingSpeedLimit(form.speedId);
|
||||
const normalizedIPSpeedId = normalizeSpeedId(form.ipSpeedId);
|
||||
const ipSpeedLimitAutoCleared = isMissingSpeedLimit(form.ipSpeedId);
|
||||
|
||||
if (isEdit) {
|
||||
const updateData = {
|
||||
@@ -2256,6 +2280,8 @@ export default function ForwardPage() {
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizedSpeedId,
|
||||
ipMaxConn: form.ipMaxConn,
|
||||
...(isAdmin ? { ipSpeedId: normalizedIPSpeedId } : {}),
|
||||
maxConn: form.maxConn,
|
||||
proxyProtocol: form.proxyProtocol,
|
||||
};
|
||||
@@ -2270,6 +2296,8 @@ export default function ForwardPage() {
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizedSpeedId,
|
||||
ipMaxConn: form.ipMaxConn,
|
||||
...(isAdmin ? { ipSpeedId: normalizedIPSpeedId } : {}),
|
||||
maxConn: form.maxConn,
|
||||
proxyProtocol: form.proxyProtocol,
|
||||
};
|
||||
@@ -2297,6 +2325,12 @@ export default function ForwardPage() {
|
||||
duration: 5000,
|
||||
});
|
||||
}
|
||||
if (isAdmin && ipSpeedLimitAutoCleared) {
|
||||
toast("所选每 IP 限速规则不存在,已自动清除为不限速", {
|
||||
icon: "⚠️",
|
||||
duration: 5000,
|
||||
});
|
||||
}
|
||||
toast.success(isEdit ? "修改成功" : "创建成功");
|
||||
setModalOpen(false);
|
||||
await refreshForwardList(false);
|
||||
@@ -2600,7 +2634,7 @@ export default function ForwardPage() {
|
||||
try {
|
||||
document.execCommand("copy");
|
||||
toast.success(`已复制${label}`);
|
||||
} catch (err) {
|
||||
} catch {
|
||||
toast.error("复制失败");
|
||||
}
|
||||
document.body.removeChild(textArea);
|
||||
@@ -4884,6 +4918,7 @@ export default function ForwardPage() {
|
||||
<AccordionItem
|
||||
key="advanced"
|
||||
aria-label="高级设置"
|
||||
className="border-b-0 [&_[data-slot=accordion-trigger]]:no-underline [&_[data-slot=accordion-trigger]]:hover:no-underline"
|
||||
title={
|
||||
<span className="text-small text-default-500 font-medium">
|
||||
高级设置
|
||||
@@ -4892,10 +4927,10 @@ export default function ForwardPage() {
|
||||
>
|
||||
<div className="space-y-4 pb-2">
|
||||
<Input
|
||||
description="此设置优先于用户的全局连接数限制。0 表示不限制。"
|
||||
description="大于 0 时优先于用户全局限制;0 或空表示使用用户全局限制,用户也为 0 时不限制。"
|
||||
label="最大连接数"
|
||||
min="0"
|
||||
placeholder="0 或空表示不限制"
|
||||
placeholder="0 或空表示使用用户全局限制"
|
||||
type="number"
|
||||
value={
|
||||
form.maxConn === 0 ? "" : String(form.maxConn || "")
|
||||
@@ -4910,6 +4945,27 @@ export default function ForwardPage() {
|
||||
setForm((prev) => ({ ...prev, maxConn: value }));
|
||||
}}
|
||||
/>
|
||||
<Input
|
||||
description="每个客户端 IP 可同时建立的最大连接数;0 或空表示不限制。"
|
||||
label="每 IP 最大连接数"
|
||||
min="0"
|
||||
placeholder="0 或空表示不限制"
|
||||
type="number"
|
||||
value={
|
||||
form.ipMaxConn === 0
|
||||
? ""
|
||||
: String(form.ipMaxConn || "")
|
||||
}
|
||||
variant="bordered"
|
||||
onChange={(e) => {
|
||||
const value = Math.max(
|
||||
Number(e.target.value) || 0,
|
||||
0,
|
||||
);
|
||||
|
||||
setForm((prev) => ({ ...prev, ipMaxConn: value }));
|
||||
}}
|
||||
/>
|
||||
<Select
|
||||
description="启用 PROXY protocol,用于透传客户端真实 IP"
|
||||
label="Proxy Protocol"
|
||||
@@ -4962,6 +5018,40 @@ export default function ForwardPage() {
|
||||
))}
|
||||
</Select>
|
||||
)}
|
||||
{isAdmin && (
|
||||
<Select
|
||||
description="每个客户端 IP 独享该限速规则;不选择表示不限制。"
|
||||
label="每 IP 限速"
|
||||
placeholder="不限速"
|
||||
selectedKeys={
|
||||
selectedIPSpeedId !== null
|
||||
? [selectedIPSpeedId.toString()]
|
||||
: []
|
||||
}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as
|
||||
| string
|
||||
| undefined;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
ipSpeedId: selectedKey
|
||||
? Number(selectedKey)
|
||||
: null,
|
||||
}));
|
||||
}}
|
||||
>
|
||||
{availableSpeedLimits.map((speedLimit) => (
|
||||
<SelectItem
|
||||
key={speedLimit.id.toString()}
|
||||
textValue={speedLimit.name}
|
||||
>
|
||||
{speedLimit.name}
|
||||
</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
)}
|
||||
</div>
|
||||
</AccordionItem>
|
||||
</Accordion>
|
||||
|
||||
@@ -10,7 +10,6 @@
|
||||
*/
|
||||
|
||||
import { registerTheme } from "./registry";
|
||||
|
||||
// ── Built-in themes ──────────────────────────────────────────────────────────
|
||||
import defaultTheme from "./default";
|
||||
import cyberpunkTheme from "./example-cyberpunk";
|
||||
|
||||
@@ -117,8 +117,6 @@ export function activateTheme(id: string): void {
|
||||
const pkg = installed.get(id);
|
||||
|
||||
if (!pkg) {
|
||||
console.warn(`[FLVX themes] Theme "${id}" is not registered.`);
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user