mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 15:46:38 +08:00
Compare commits
39 Commits
3.0.1-alpha1
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 129fa0aa4c | |||
| cf7246c71b | |||
| b4c2989285 | |||
| f2783713a5 | |||
| 5d60c4fbe1 | |||
| 62c56e0a79 | |||
| 2e845030de | |||
| 2269f2e2d5 | |||
| da7bef88f1 | |||
| f26014579b | |||
| 5041c722c9 | |||
| c56798e991 | |||
| 40e96f3592 | |||
| 0e24b53a5b | |||
| a820c49c94 | |||
| 538e64ffc0 | |||
| 0b23d6f7d7 | |||
| 9e6f80019d | |||
| a8fd01d4d8 | |||
| cbe2fc492e | |||
| ae370382d3 | |||
| e112d81697 | |||
| 11a27d3c67 | |||
| 8e513a1bae | |||
| f98be845d3 | |||
| 0c2acfdd8a | |||
| 777db8767f | |||
| 82f6047506 | |||
| 3ce320da5a | |||
| 7ab0db29ae | |||
| 006ea97200 | |||
| e569aedd3e | |||
| 35080aea2d | |||
| 85e588ffe9 | |||
| 03524f4a65 | |||
| 9e69e020ab | |||
| 14bbd3907d | |||
| ca8d8e92ba | |||
| 079474fa06 |
@@ -7,6 +7,18 @@ on:
|
||||
branches: ['**']
|
||||
|
||||
jobs:
|
||||
install-scripts:
|
||||
name: Test Installer Scripts (systemd and Alpine OpenRC)
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Run installer regression tests
|
||||
run: bash test-install-scripts-proxy.sh
|
||||
|
||||
- name: Test Alpine bootstrap and OpenRC lifecycle
|
||||
run: docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh
|
||||
|
||||
frontend:
|
||||
name: Build Frontend
|
||||
runs-on: ubuntu-latest
|
||||
@@ -22,7 +34,7 @@ jobs:
|
||||
node-version: '20.19.0'
|
||||
|
||||
- name: Install pnpm
|
||||
run: npm install -g pnpm
|
||||
run: npm install -g pnpm@10.28.1
|
||||
|
||||
- name: Install dependencies
|
||||
run: pnpm install --frozen-lockfile
|
||||
|
||||
@@ -229,6 +229,7 @@ jobs:
|
||||
|
||||
docker buildx build \
|
||||
--platform linux/amd64,linux/arm64 \
|
||||
--build-arg KEYGEN_ACCOUNT_ID=${{ secrets.KEYGEN_ACCOUNT_ID }} \
|
||||
--push \
|
||||
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:latest \
|
||||
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:${VERSION} \
|
||||
@@ -303,6 +304,9 @@ jobs:
|
||||
sed -i "s|^PINNED_VERSION=\"\"|PINNED_VERSION=\"${VERSION}\"|" ./artifacts/install.sh
|
||||
sed -i "s|^PINNED_VERSION=\"\"|PINNED_VERSION=\"${VERSION}\"|" ./artifacts/panel_install.sh
|
||||
|
||||
- name: Verify release installer on Alpine OpenRC
|
||||
run: docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh /workspace/artifacts/install.sh
|
||||
|
||||
- name: Create Release
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
@@ -431,4 +435,3 @@ jobs:
|
||||
gh release upload "${VERSION}" ./artifacts/gost-arm64.sha256 --clobber
|
||||
|
||||
echo "✅ GOST 二进制文件更新完成"
|
||||
|
||||
|
||||
+10
-1
@@ -64,6 +64,14 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/panel_instal
|
||||
curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -o install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
Alpine Linux 最小化安装若未包含 `curl`,可使用系统自带的 `wget` 下载:
|
||||
|
||||
```bash
|
||||
wget -O install.sh https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh && chmod +x install.sh && ./install.sh
|
||||
```
|
||||
|
||||
脚本会在 Alpine 上自动安装 Bash、`curl` 和 CA 证书,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
|
||||
|
||||
**安装过程中会提示输入:**
|
||||
- **服务器地址**: 面板端的通信地址(通常是 `http://<面板IP>:<后端端口>`,例如 `http://1.2.3.4:6365`)。
|
||||
- **密钥**: 刚才在面板中获取的节点密钥。
|
||||
@@ -77,7 +85,8 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -
|
||||
|
||||
### 3. 验证安装
|
||||
安装完成后,服务会自动启动。
|
||||
- 查看状态: `systemctl status flux_agent`
|
||||
- systemd 查看状态: `systemctl status flux_agent`
|
||||
- Alpine/OpenRC 查看状态: `rc-service flux_agent status`
|
||||
- 回到面板 **节点管理** 页面,该节点状态应显示为 **在线**。
|
||||
|
||||
---
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,650 @@
|
||||
# Forward Flow Reset 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 permission-checked action that resets only one forward rule's displayed upload and download counters.
|
||||
|
||||
**Architecture:** A dedicated repository method updates only the selected `forward` row. A dedicated authenticated handler reuses `resolveForwardAccess`, and the React page calls the endpoint from all three rule views through one confirmation modal.
|
||||
|
||||
**Tech Stack:** Go `net/http`, GORM, SQLite/PostgreSQL-compatible models, React, TypeScript, shadcn bridge components, Tailwind CSS v4.
|
||||
|
||||
## Global Constraints
|
||||
|
||||
- Only `forward.in_flow`, `forward.out_flow`, and `forward.updated_time` may change during reset.
|
||||
- Do not modify `user`, `user_tunnel`, quota, historical statistics, nftables counter state, or running services.
|
||||
- Administrators may reset any rule; non-admin users may reset only their own rules through existing `resolveForwardAccess` behavior.
|
||||
- All API responses must keep the `{code, msg, data, ts}` envelope.
|
||||
- Frontend imports must use `src/shadcn-bridge/heroui/*`; do not add `@heroui/*` or `@nextui-org/*` dependencies.
|
||||
- Do not add frontend test infrastructure.
|
||||
- Do not edit generated protobuf files, `install.sh`, or `panel_install.sh`.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Add the repository flow-reset primitive
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/store/repo/repository_forward_flow_reset_test.go`
|
||||
- Modify: `go-backend/internal/store/repo/repository_mutations.go`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `model.Forward`, the repository's GORM database handle, and an explicit Unix-millisecond timestamp.
|
||||
- Produces: `func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error`.
|
||||
|
||||
- [ ] **Step 1: Write the failing repository tests**
|
||||
|
||||
Create `go-backend/internal/store/repo/repository_forward_flow_reset_test.go`:
|
||||
|
||||
```go
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResetForwardFlowOnlyUpdatesSelectedForward(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "forward-flow-reset.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
const originalUpdated int64 = 1000
|
||||
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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)
|
||||
`).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, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert 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(10, 2, 1, 10, 100, 500, 600, 0, 0, 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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, ?, 1, 0),
|
||||
(21, 2, 'owner', 'other', 1, '127.0.0.1:81', 'fifo', 333, 444, 1000, ?, 1, 1)
|
||||
`, originalUpdated, originalUpdated).Error; err != nil {
|
||||
t.Fatalf("insert forwards: %v", err)
|
||||
}
|
||||
|
||||
const resetAt int64 = 2000
|
||||
if err := r.ResetForwardFlow(20, resetAt); err != nil {
|
||||
t.Fatalf("ResetForwardFlow: %v", err)
|
||||
}
|
||||
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", 0)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", 0)
|
||||
assertForwardFlowResetValue(t, r, "SELECT updated_time FROM forward WHERE id = 20", resetAt)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM forward WHERE id = 21", 333)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM forward WHERE id = 21", 444)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM user WHERE id = 2", 900)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM user_tunnel WHERE id = 10", 500)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
}
|
||||
|
||||
func TestResetForwardFlowRejectsUninitializedRepository(t *testing.T) {
|
||||
var r *Repository
|
||||
if err := r.ResetForwardFlow(20, 2000); err == nil {
|
||||
t.Fatal("expected uninitialized repository error")
|
||||
}
|
||||
}
|
||||
|
||||
func assertForwardFlowResetValue(t *testing.T, r *Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run the repository tests and verify the missing method failure**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./internal/store/repo -run TestResetForwardFlow -count=1
|
||||
```
|
||||
|
||||
Expected: compilation fails because `ResetForwardFlow` is undefined.
|
||||
|
||||
- [ ] **Step 3: Implement the minimal repository method**
|
||||
|
||||
Add to the flow-reset section of `go-backend/internal/store/repo/repository_mutations.go`:
|
||||
|
||||
```go
|
||||
func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", forwardID).
|
||||
Updates(map[string]interface{}{
|
||||
"in_flow": 0,
|
||||
"out_flow": 0,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
```
|
||||
|
||||
The file already imports `errors` and `model`; do not add a new dependency.
|
||||
|
||||
- [ ] **Step 4: Format and run the focused repository tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && gofmt -w internal/store/repo/repository_forward_flow_reset_test.go internal/store/repo/repository_mutations.go
|
||||
go test ./internal/store/repo -run TestResetForwardFlow -count=1
|
||||
```
|
||||
|
||||
Expected: both reset tests pass.
|
||||
|
||||
- [ ] **Step 5: Commit the repository change**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/store/repo/repository_mutations.go go-backend/internal/store/repo/repository_forward_flow_reset_test.go
|
||||
git commit -m "feat: add forward flow reset repository method"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 2: Add the authenticated reset endpoint
|
||||
|
||||
**Files:**
|
||||
- Create: `go-backend/internal/http/handler/forward_reset_flow_test.go`
|
||||
- Modify: `go-backend/internal/http/handler/handler.go`
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `POST` JSON `{ "id": number }`, `resolveForwardAccess`, and `Repository.ResetForwardFlow` from Task 1.
|
||||
- Produces: `POST /api/v1/forward/reset-flow` and `func (h *Handler) forwardResetFlow(http.ResponseWriter, *http.Request)`.
|
||||
|
||||
- [ ] **Step 1: Write the failing handler tests**
|
||||
|
||||
Create `go-backend/internal/http/handler/forward_reset_flow_test.go`:
|
||||
|
||||
```go
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestForwardResetFlowPermissionsAndIsolation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
actorID int64
|
||||
actorRole int
|
||||
forwardID int64
|
||||
wantCode int
|
||||
wantInFlow int64
|
||||
wantOutFlow int64
|
||||
}{
|
||||
{name: "admin resets another user's rule", actorID: 1, actorRole: 0, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "owner resets own rule", actorID: 2, actorRole: 1, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "user cannot reset another user's rule", actorID: 3, actorRole: 1, forwardID: 20, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
{name: "missing rule is rejected", actorID: 1, actorRole: 0, forwardID: 999, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
h, r := setupForwardResetFlowHandler(t)
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, tt.forwardID, tt.actorID, tt.actorRole)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
h.forwardResetFlow(res, req)
|
||||
|
||||
if got := decodeForwardResetFlowCode(t, res); got != tt.wantCode {
|
||||
t.Fatalf("code = %d, want %d; body=%s", got, tt.wantCode, res.Body.String())
|
||||
}
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", tt.wantInFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", tt.wantOutFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResetFlowRejectsInvalidRequests(t *testing.T) {
|
||||
h, _ := setupForwardResetFlowHandler(t)
|
||||
|
||||
t.Run("non post", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodGet, 20, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid id", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, 0, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func setupForwardResetFlowHandler(t *testing.T) (*Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "forward-reset-handler.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
statements := []string{
|
||||
`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(1, 'admin', 'pwd', 0, 0, 100, 0, 0, 0, 10, 1000, 1000, 1)`,
|
||||
`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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)`,
|
||||
`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, 'other', 'pwd', 1, 0, 100, 0, 0, 0, 10, 1000, 1000, 1)`,
|
||||
`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)`,
|
||||
`INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, 10, 100, 500, 600, 0, 0, 1)`,
|
||||
`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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, 1000, 1, 0)`,
|
||||
}
|
||||
for _, statement := range statements {
|
||||
if err := r.DB().Exec(statement).Error; err != nil {
|
||||
t.Fatalf("seed database: %v", err)
|
||||
}
|
||||
}
|
||||
return New(r, "test-secret"), r
|
||||
}
|
||||
|
||||
func newForwardResetFlowRequest(t *testing.T, method string, forwardID, actorID int64, roleID int) *http.Request {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(map[string]int64{"id": forwardID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(method, "/api/v1/forward/reset-flow", bytes.NewReader(body))
|
||||
claims := auth.Claims{Sub: strconv.FormatInt(actorID, 10), RoleID: roleID}
|
||||
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
|
||||
}
|
||||
|
||||
func decodeForwardResetFlowCode(t *testing.T, res *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
}
|
||||
if err := json.Unmarshal(res.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, res.Body.String())
|
||||
}
|
||||
return payload.Code
|
||||
}
|
||||
|
||||
func assertForwardResetFlowDBValue(t *testing.T, r *repo.Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
If the project's default error code differs from `-1`, replace the test expectation with the actual `response.ErrDefault` code after inspecting one existing handler response; do not weaken the success and database assertions.
|
||||
|
||||
- [ ] **Step 2: Run the handler tests and verify the missing handler failure**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./internal/http/handler -run TestForwardResetFlow -count=1
|
||||
```
|
||||
|
||||
Expected: compilation fails because `forwardResetFlow` is undefined.
|
||||
|
||||
- [ ] **Step 3: Register and implement the endpoint**
|
||||
|
||||
Add this route beside the other forward routes in `go-backend/internal/http/handler/handler.go`:
|
||||
|
||||
```go
|
||||
mux.HandleFunc("/api/v1/forward/reset-flow", h.forwardResetFlow)
|
||||
```
|
||||
|
||||
Add this handler beside `forwardPause` and `forwardResume` in `go-backend/internal/http/handler/mutations.go`:
|
||||
|
||||
```go
|
||||
func (h *Handler) forwardResetFlow(w http.ResponseWriter, r *http.Request) {
|
||||
id := idFromBody(r, w)
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
if _, _, _, err := h.resolveForwardAccess(r, id); err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.ResetForwardFlow(id, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
```
|
||||
|
||||
This deliberately does not call runtime service controls or nftables reconciliation.
|
||||
|
||||
- [ ] **Step 4: Format and run the focused handler tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && gofmt -w internal/http/handler/forward_reset_flow_test.go internal/http/handler/handler.go internal/http/handler/mutations.go
|
||||
go test ./internal/http/handler -run TestForwardResetFlow -count=1
|
||||
```
|
||||
|
||||
Expected: all reset endpoint tests pass.
|
||||
|
||||
- [ ] **Step 5: Run all backend tests**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./...
|
||||
```
|
||||
|
||||
Expected: all backend packages and contract tests pass, excluding environment-gated PostgreSQL tests when `FLVX_POSTGRES_TEST_DSN` is unset.
|
||||
|
||||
- [ ] **Step 6: Commit the endpoint change**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/handler.go go-backend/internal/http/handler/mutations.go go-backend/internal/http/handler/forward_reset_flow_test.go
|
||||
git commit -m "feat: add forward flow reset endpoint"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 3: Add the rule-page reset action and confirmation modal
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/api/index.ts`
|
||||
- Modify: `vite-frontend/src/pages/forward.tsx`
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: `POST /forward/reset-flow`, the page's `Forward` shape, `refreshForwardList`, toast notifications, and existing modal/button bridge components.
|
||||
- Produces: `resetForwardFlow(id: number)`, a shared reset handler, disabled zero-usage actions in all rule views, and one confirmation modal.
|
||||
|
||||
- [ ] **Step 1: Add the frontend API wrapper**
|
||||
|
||||
Add beside the forward control operations in `vite-frontend/src/api/index.ts`:
|
||||
|
||||
```ts
|
||||
export const resetForwardFlow = (forwardId: number) =>
|
||||
Network.post("/forward/reset-flow", { id: forwardId });
|
||||
```
|
||||
|
||||
Import `resetForwardFlow` from `@/api` in `vite-frontend/src/pages/forward.tsx`.
|
||||
|
||||
- [ ] **Step 2: Add page state and shared reset handlers**
|
||||
|
||||
Add state beside the existing delete modal state:
|
||||
|
||||
```ts
|
||||
const [resetFlowModalOpen, setResetFlowModalOpen] = useState(false);
|
||||
const [resetFlowLoading, setResetFlowLoading] = useState(false);
|
||||
const [forwardToResetFlow, setForwardToResetFlow] = useState<Forward | null>(null);
|
||||
```
|
||||
|
||||
Add these handlers beside `handleDelete` and `confirmDelete`:
|
||||
|
||||
```ts
|
||||
const handleResetFlow = (forward: Forward) => {
|
||||
if ((forward.inFlow || 0) + (forward.outFlow || 0) <= 0) return;
|
||||
setForwardToResetFlow(forward);
|
||||
setResetFlowModalOpen(true);
|
||||
};
|
||||
|
||||
const confirmResetFlow = async () => {
|
||||
if (!forwardToResetFlow) return;
|
||||
|
||||
setResetFlowLoading(true);
|
||||
try {
|
||||
const res = await resetForwardFlow(forwardToResetFlow.id);
|
||||
|
||||
if (res.code !== 0) {
|
||||
toast.error(res.msg || "流量清零失败");
|
||||
return;
|
||||
}
|
||||
|
||||
toast.success("规则流量已清零");
|
||||
setResetFlowModalOpen(false);
|
||||
setForwardToResetFlow(null);
|
||||
await refreshForwardList(false);
|
||||
} catch {
|
||||
toast.error("流量清零失败");
|
||||
} finally {
|
||||
setResetFlowLoading(false);
|
||||
}
|
||||
};
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Add one reusable reset icon button to both table row components**
|
||||
|
||||
Pass `handleResetFlow` into `SortableTableRow` and `SortableCompactTableRow` at every render site. Add it to each component's destructured props.
|
||||
|
||||
Insert this button between diagnosis and delete in each table action cell:
|
||||
|
||||
```tsx
|
||||
<Button
|
||||
isIconOnly
|
||||
className="bg-secondary/10 text-secondary hover:bg-secondary/20"
|
||||
isDisabled={(forward.inFlow || 0) + (forward.outFlow || 0) <= 0}
|
||||
size="sm"
|
||||
title="流量清零"
|
||||
onPress={() => handleResetFlow(forward)}
|
||||
>
|
||||
<svg
|
||||
aria-hidden="true"
|
||||
className="h-4 w-4"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
d="M4 4v6h6M20 20v-6h-6M20 9a8 8 0 00-13.657-3.657L4 8m16 8-2.343 2.657A8 8 0 014 15"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
/>
|
||||
</svg>
|
||||
</Button>
|
||||
```
|
||||
|
||||
- [ ] **Step 4: Add the reset action to the card view**
|
||||
|
||||
Insert a fourth action button between diagnosis and delete in `renderForwardCard`:
|
||||
|
||||
```tsx
|
||||
<Button
|
||||
className="flex-1 min-h-8"
|
||||
color="secondary"
|
||||
isDisabled={(forward.inFlow || 0) + (forward.outFlow || 0) <= 0}
|
||||
size="sm"
|
||||
startContent={
|
||||
<svg
|
||||
aria-hidden="true"
|
||||
className="w-3 h-3"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
d="M4 4v6h6M20 20v-6h-6M20 9a8 8 0 00-13.657-3.657L4 8m16 8-2.343 2.657A8 8 0 014 15"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
/>
|
||||
</svg>
|
||||
}
|
||||
variant="flat"
|
||||
onPress={() => handleResetFlow(forward)}
|
||||
>
|
||||
清零
|
||||
</Button>
|
||||
```
|
||||
|
||||
Change the card action container from `flex gap-1.5 mt-3` to `grid grid-cols-2 gap-1.5 mt-3` so all four actions remain readable at the smallest supported card width.
|
||||
|
||||
- [ ] **Step 5: Add the confirmation modal**
|
||||
|
||||
Add beside the delete confirmation modal:
|
||||
|
||||
```tsx
|
||||
<Modal
|
||||
backdrop="blur"
|
||||
classNames={{
|
||||
base: "!w-[calc(100%-32px)] !mx-auto sm:!w-full rounded-2xl overflow-hidden",
|
||||
}}
|
||||
isOpen={resetFlowModalOpen}
|
||||
placement="center"
|
||||
scrollBehavior="inside"
|
||||
size="lg"
|
||||
onOpenChange={setResetFlowModalOpen}
|
||||
>
|
||||
<ModalContent>
|
||||
{(onClose) => (
|
||||
<>
|
||||
<ModalHeader className="flex flex-col gap-1">
|
||||
<h2 className="text-lg font-bold text-secondary">确认流量清零</h2>
|
||||
</ModalHeader>
|
||||
<ModalBody>
|
||||
<p className="text-default-600">
|
||||
确定要清零规则{" "}
|
||||
<span className="font-semibold text-foreground">
|
||||
"{forwardToResetFlow?.name}"
|
||||
</span>{" "}
|
||||
当前显示的上传和下载流量吗?
|
||||
</p>
|
||||
<p className="text-small text-default-500 mt-2">
|
||||
此操作不可撤销,但不会影响用户总流量、用户隧道配额和历史统计。
|
||||
</p>
|
||||
</ModalBody>
|
||||
<ModalFooter>
|
||||
<Button isDisabled={resetFlowLoading} variant="light" onPress={onClose}>
|
||||
取消
|
||||
</Button>
|
||||
<Button
|
||||
color="secondary"
|
||||
isLoading={resetFlowLoading}
|
||||
onPress={confirmResetFlow}
|
||||
>
|
||||
确认清零
|
||||
</Button>
|
||||
</ModalFooter>
|
||||
</>
|
||||
)}
|
||||
</ModalContent>
|
||||
</Modal>
|
||||
```
|
||||
|
||||
Add this wrapper beside the other reset handlers and pass it to the modal as `onOpenChange={handleResetFlowModalOpenChange}`:
|
||||
|
||||
```ts
|
||||
const handleResetFlowModalOpenChange = (isOpen: boolean) => {
|
||||
if (resetFlowLoading) return;
|
||||
setResetFlowModalOpen(isOpen);
|
||||
if (!isOpen) {
|
||||
setForwardToResetFlow(null);
|
||||
}
|
||||
};
|
||||
```
|
||||
|
||||
- [ ] **Step 6: Format and verify the frontend**
|
||||
|
||||
Run:
|
||||
|
||||
```bash
|
||||
cd vite-frontend && pnpm exec prettier --write src/api/index.ts src/pages/forward.tsx
|
||||
pnpm run build
|
||||
pnpm run lint
|
||||
```
|
||||
|
||||
Expected: TypeScript/Vite build succeeds and ESLint finishes without errors.
|
||||
|
||||
- [ ] **Step 7: Commit the frontend change**
|
||||
|
||||
```bash
|
||||
git add vite-frontend/src/api/index.ts vite-frontend/src/pages/forward.tsx
|
||||
git commit -m "feat: add forward flow reset action"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### Task 4: Perform integrated verification
|
||||
|
||||
**Files:**
|
||||
- Verify only; no planned source changes.
|
||||
|
||||
**Interfaces:**
|
||||
- Consumes: the repository method, API endpoint, and rule-page action from Tasks 1-3.
|
||||
- Produces: evidence that the complete feature builds and all affected tests pass.
|
||||
|
||||
- [ ] **Step 1: Run the complete backend suite**
|
||||
|
||||
```bash
|
||||
cd go-backend && go test ./...
|
||||
```
|
||||
|
||||
Expected: all available backend tests pass.
|
||||
|
||||
- [ ] **Step 2: Run the complete frontend checks**
|
||||
|
||||
```bash
|
||||
cd vite-frontend && pnpm run build && pnpm run lint
|
||||
```
|
||||
|
||||
Expected: both commands exit successfully.
|
||||
|
||||
- [ ] **Step 3: Check formatting and working-tree scope**
|
||||
|
||||
```bash
|
||||
git diff --check
|
||||
git status --short
|
||||
git log -4 --oneline
|
||||
```
|
||||
|
||||
Expected: no whitespace errors; the working tree is clean; the three feature commits are visible after the design and implementation-plan commits.
|
||||
|
||||
- [ ] **Step 4: Manually verify the feature when a local panel is available**
|
||||
|
||||
1. Open the Rules page as an administrator and reset a rule with non-zero upload/download traffic.
|
||||
2. Confirm the modal states that user totals, tunnel quota, and history are unaffected.
|
||||
3. Confirm the rule immediately shows zero after success.
|
||||
4. Confirm the user page's total traffic and user-tunnel traffic values did not change.
|
||||
5. Generate new traffic and confirm the rule starts accumulating from zero.
|
||||
6. Log in as a normal user and confirm the user can reset an owned rule but cannot access another user's rule through a direct API request.
|
||||
|
||||
Expected: all six checks match the design specification.
|
||||
@@ -0,0 +1,284 @@
|
||||
# nftables 流量统计设计
|
||||
|
||||
**日期**: 2026-06-06
|
||||
**状态**: 待审核
|
||||
**作者**: Codex
|
||||
|
||||
## 概述
|
||||
|
||||
为 FLVX 的 `nftables` 转发模式补齐流量统计。当前 nftables 模式由面板通过 SSH 全量维护 `table inet flvx`,但没有 agent,因此不能复用 WebSocket 运行时上报。新方案由面板定时通过 SSH 拉取远端 nftables counter,计算增量后写入现有流量账本。
|
||||
|
||||
目标是让 nftables 转发在用户可见口径上尽量接近 agent 模式:
|
||||
|
||||
- forward 列表显示 `inFlow` / `outFlow`。
|
||||
- 用户、用户隧道、配额和流量策略继续生效。
|
||||
- 隧道监控继续获得分钟级 `tunnel_metric`。
|
||||
- 节点不需要安装新的 agent 或常驻进程。
|
||||
|
||||
## 背景
|
||||
|
||||
现有 agent 模式通过 `/flow/upload` 接收加密上报,handler 会把服务名解析为 `forward_id/user_id/user_tunnel_id`,再复用以下路径:
|
||||
|
||||
- `ApplyFlowUploadDeltasBatch` 更新 `forward`、`user`、`user_tunnel`。
|
||||
- `AddUserQuotaUsageBatch` 更新用户配额窗口。
|
||||
- `enforceUserQuotaIfNeeded` 和 `enforceFlowPolicies` 做约束 enforcement。
|
||||
- `recordTunnelMetricsFromForwardBatch` 写入分钟级隧道监控。
|
||||
|
||||
nftables 模式已经有 `nft_rule_binding` 记录规则应用状态,规则 comment 里包含 `forward_id`。这给 counter 到业务实体的映射提供了稳定锚点。
|
||||
|
||||
## 推荐方案
|
||||
|
||||
采用“面板 SSH 轮询 nftables counter”的方案:
|
||||
|
||||
1. 渲染 nftables 规则时,为每个 forward、协议和方向写入稳定 comment 和 `counter`。
|
||||
2. 后端定时扫描 `forward_mode = nftables` 的节点。
|
||||
3. 对每个节点通过 SSH 执行 `nft -j list table inet flvx`。
|
||||
4. 解析 JSON 规则,按 comment 得到 `forward_id/protocol/direction/bytes/packets`。
|
||||
5. 用数据库中的上次采样值计算 delta。
|
||||
6. 将 delta 转成现有 flow upload 内部结构,复用既有入账、配额、策略和监控逻辑。
|
||||
|
||||
不采用节点 crontab 或 systemd timer 回推。它会重新引入节点侧组件,削弱 nftables 模式“不安装 agent”的产品边界。
|
||||
|
||||
## 统计口径
|
||||
|
||||
正式入账使用 `forward` filter chain 的计数,不使用 NAT chain 的 DNAT 命中计数作为主口径。
|
||||
|
||||
原因:
|
||||
|
||||
- DNAT counter 表示规则命中,不一定代表后续转发成功。
|
||||
- filter forward chain 更接近实际经过内核转发的数据。
|
||||
- SNAT/masquerade 会改变包头,入账规则应在可稳定匹配目标服务地址和端口的位置统计。
|
||||
|
||||
方向定义:
|
||||
|
||||
| direction | nft 匹配 | 写入字段 |
|
||||
|-----------|----------|----------|
|
||||
| `to-target` | 外部客户端到目标服务 | `in_flow` |
|
||||
| `from-target` | 目标服务返回外部客户端 | `out_flow` |
|
||||
|
||||
用户总用量和配额仍按 `in_flow + out_flow` 计算。隧道 `traffic_ratio` 和 `flow` 倍率继续沿用 agent 模式逻辑,保证不同运行时模式的账单口径一致。
|
||||
|
||||
## nftables 规则设计
|
||||
|
||||
继续只维护 `table inet flvx`,避免触碰用户已有规则。每条 forward 对 TCP 和 UDP 各生成一组 DNAT 和统计规则。
|
||||
|
||||
示例:
|
||||
|
||||
```nft
|
||||
table inet flvx {
|
||||
chain prerouting {
|
||||
type nat hook prerouting priority dstnat; policy accept;
|
||||
tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"
|
||||
udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"
|
||||
}
|
||||
|
||||
chain postrouting {
|
||||
type nat hook postrouting priority srcnat; policy accept;
|
||||
masquerade comment "flvx masquerade"
|
||||
}
|
||||
|
||||
chain forward {
|
||||
type filter hook forward priority filter; policy accept;
|
||||
ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"
|
||||
ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"
|
||||
ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"
|
||||
ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
IPv6 目标使用 `ip6`:
|
||||
|
||||
```nft
|
||||
ip6 daddr 2001:db8::20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"
|
||||
ip6 saddr 2001:db8::20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"
|
||||
```
|
||||
|
||||
域名目标无法在 nftables 规则中动态匹配返回方向。统计第一阶段要求 nftables forward 的 `remoteAddr` host 必须是 IP 地址;如果当前纯转发实现允许域名,开启统计时应同步收紧校验。后续若要支持域名,应在规则同步时解析并固化 IP,同时明确 DNS 变化后的重建策略。
|
||||
|
||||
## Comment 格式
|
||||
|
||||
正式统计规则使用固定格式:
|
||||
|
||||
```text
|
||||
flvx forward:<forward_id> <direction> <protocol>
|
||||
```
|
||||
|
||||
字段:
|
||||
|
||||
- `forward_id`: 十进制整数。
|
||||
- `direction`: `to-target` 或 `from-target`。
|
||||
- `protocol`: `tcp` 或 `udp`。
|
||||
|
||||
DNAT 调试规则可使用 `dnat` direction,但 collector 不入账 `dnat`。后端只依赖 comment 解析,不依赖 nft handle,因为全量重建 table 会改变 handle。
|
||||
|
||||
## 数据模型
|
||||
|
||||
新增 `nft_counter_state` 表保存上次采样基线。
|
||||
|
||||
| 字段 | 说明 |
|
||||
|------|------|
|
||||
| `id` | 主键 |
|
||||
| `node_id` | nftables 节点 ID |
|
||||
| `forward_id` | 转发规则 ID |
|
||||
| `protocol` | `tcp` / `udp` |
|
||||
| `direction` | `to-target` / `from-target` |
|
||||
| `rule_hash` | 当前规则 hash |
|
||||
| `bytes` | 上次采样绝对字节数 |
|
||||
| `packets` | 上次采样绝对包数 |
|
||||
| `collected_time` | 上次采样时间 |
|
||||
| `created_time` | 创建时间 |
|
||||
| `updated_time` | 更新时间 |
|
||||
|
||||
唯一索引:
|
||||
|
||||
```text
|
||||
node_id, forward_id, protocol, direction
|
||||
```
|
||||
|
||||
GORM 模型必须定义 `TableName()`,字段 tag 保持 SQLite/PostgreSQL 兼容,不使用 `jsonb`、`serial` 等数据库专属类型。
|
||||
|
||||
## 后端组件
|
||||
|
||||
扩展 `go-backend/internal/runtime/nftables`:
|
||||
|
||||
| 组件 | 职责 |
|
||||
|------|------|
|
||||
| `CounterSample` | 表达单条 nft counter 采样 |
|
||||
| `Collector` | 对外提供 `Collect(ctx, cfg)` |
|
||||
| `SSHRunner.ListTableJSON` | 远端执行 `nft -j list table inet flvx` |
|
||||
| `ParseCounterSamples` | 解析 nft JSON 和 FLVX comment |
|
||||
|
||||
扩展 repository:
|
||||
|
||||
| 方法 | 职责 |
|
||||
|------|------|
|
||||
| `ListNftablesNodesForCollection` | 找到启用 nftables 且有 SSH 配置的节点 |
|
||||
| `GetNftCounterStatesByNode` | 读取节点上次 counter 基线 |
|
||||
| `UpsertNftCounterStates` | 批量刷新基线 |
|
||||
| `DeleteNftCounterStatesByForward` | forward 删除时清理状态 |
|
||||
|
||||
扩展 handler/job:
|
||||
|
||||
- 新增 `runNftablesTrafficCollectJob(now time.Time)`。
|
||||
- 默认每 60 秒运行一次。
|
||||
- 对节点采集设置并发上限,建议 3 到 5。
|
||||
- 单节点失败只记录日志和节点采集状态,不影响其他节点。
|
||||
|
||||
## 增量算法
|
||||
|
||||
collector 返回的是 nftables 的绝对 counter。入账前必须和上次基线做差。
|
||||
|
||||
规则:
|
||||
|
||||
- 无旧状态:只保存当前值作为基线,不入账。
|
||||
- `rule_hash` 变化:只刷新基线,不入账,避免新旧规则混算。
|
||||
- 新 bytes 大于等于旧 bytes:`delta = new - old`。
|
||||
- 新 bytes 小于旧 bytes:认为远端 table 重建、counter reset 或系统重启,只刷新基线,不入账。
|
||||
- delta 为 0:刷新采集时间,不入账。
|
||||
- 样本无法映射到有效 forward:忽略并记录 debug 日志。
|
||||
|
||||
同一 forward 的 TCP/UDP delta 要先聚合,再转换成现有账本:
|
||||
|
||||
- `to-target` bytes 聚合为原始 `bytesIn`。
|
||||
- `from-target` bytes 聚合为原始 `bytesOut`。
|
||||
- 入账时按 `traffic_ratio` 和 `tunnel.flow` 计算 scaled `InFlow` / `OutFlow`。
|
||||
- 配额使用 scaled 后的 `InFlow + OutFlow`。
|
||||
- `tunnel_metric` 使用原始 `bytesIn` / `bytesOut`。
|
||||
|
||||
## 入账路径
|
||||
|
||||
新增一个 nftables 专用的 batch builder,但输出沿用现有结构:
|
||||
|
||||
```go
|
||||
type nftTrafficDelta struct {
|
||||
ForwardID int64
|
||||
BytesIn int64
|
||||
BytesOut int64
|
||||
}
|
||||
```
|
||||
|
||||
处理流程:
|
||||
|
||||
1. 收集本轮所有 `forward_id`。
|
||||
2. 调用 `GetFlowUploadForwardMetas` 获取 `user_id/user_tunnel_id/tunnel_id/traffic_ratio/tunnel_flow`。
|
||||
3. 构造 `repo.FlowUploadCounterDelta`。
|
||||
4. 调用 `recordTunnelMetricsFromForwardBatch` 写监控。
|
||||
5. 抽出共享入账 helper,复用 `applyFlowDeltasWithFallback`、`applyQuotaUsageWithFallback`、`enforceUserQuotaIfNeeded` 和 `enforceFlowPolicies`。不要通过伪造 agent service name 去调用 agent 专用 builder。
|
||||
|
||||
不新增独立的 nftables 流量字段。`forward.in_flow/out_flow`、`user.in_flow/out_flow`、`user_tunnel.in_flow/out_flow` 仍是统一事实来源。
|
||||
|
||||
## 错误处理
|
||||
|
||||
采集错误分为三类:
|
||||
|
||||
| 类型 | 行为 |
|
||||
|------|------|
|
||||
| SSH 连接或认证失败 | 记录日志,保留下次继续采集 |
|
||||
| 远端无 `table inet flvx` | 视为规则未应用或被清理,记录 warning,不清空账本 |
|
||||
| JSON 解析失败 | 记录原始错误摘要,不入账 |
|
||||
|
||||
不要因为采集失败禁用 forward。流量统计失败和转发运行失败不是同一件事。
|
||||
|
||||
可在后续 UI 增加节点级采集状态,例如最近成功时间、最近错误。但第一步只要求后端具备日志和数据库状态即可。
|
||||
|
||||
## 与现有行为的关系
|
||||
|
||||
- agent 模式 `/flow/upload` 不变。
|
||||
- nftables 模式不新增节点侧 HTTP 回调。
|
||||
- 现有 `nft_rule_binding.rule_hash` 继续表示规则期望状态;counter state 用它判断采样是否跨规则版本。
|
||||
- `statistics_flow` 小时统计 job 不需要改,它基于用户总流量快照自然包含 nftables 入账结果。
|
||||
- 用户重置流量时不需要清空 nftables counter。重置只清业务账本;下一轮采集继续从 counter state 差值入账。
|
||||
|
||||
## 测试计划
|
||||
|
||||
后端单元测试:
|
||||
|
||||
- renderer 为 TCP/UDP、IPv4/IPv6 目标生成 `counter` 和稳定 comment。
|
||||
- comment parser 能识别合法格式,拒绝未知 direction/protocol。
|
||||
- nft JSON parser 能从 `nft -j list table` 输出中提取 bytes/packets。
|
||||
- delta 算法覆盖首次基线、正常增长、counter reset、rule_hash 变化和零增量。
|
||||
- batch builder 正确应用 `traffic_ratio` 和 `tunnel.flow`。
|
||||
|
||||
repository 测试:
|
||||
|
||||
- `nft_counter_state` 自动迁移。
|
||||
- upsert 在 SQLite 下可重复刷新。
|
||||
- forward 删除时清理 counter state。
|
||||
|
||||
handler/job 测试:
|
||||
|
||||
- 单节点采集成功会调用现有流量入账路径。
|
||||
- 单节点 SSH 失败不影响其他节点。
|
||||
- 无旧状态时不会误把历史 counter 入账。
|
||||
|
||||
验证命令:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
```
|
||||
|
||||
## 分阶段落地
|
||||
|
||||
第一阶段:
|
||||
|
||||
- 规则渲染加入 filter chain counter。
|
||||
- 实现 SSH collector、JSON parser、counter state 和后台 job。
|
||||
- 入账到现有账本和 tunnel metric。
|
||||
|
||||
第二阶段:
|
||||
|
||||
- UI 展示 nftables 采集状态。
|
||||
- 节点详情显示最近采集时间和最近错误。
|
||||
- 提供手动“采集一次”诊断按钮。
|
||||
|
||||
第三阶段:
|
||||
|
||||
- 探索域名目标的解析和重建策略。
|
||||
- 优化大量节点下的采集调度、退避和超时配置。
|
||||
|
||||
## 开放问题
|
||||
|
||||
- 采集周期默认 60 秒是否满足产品预期;如果需要更实时,可以降到 30 秒,但 SSH 压力会增加。
|
||||
- nftables 模式是否继续允许域名 remoteAddr。如果允许,需要先定义 DNS 固化和统计匹配规则。
|
||||
- 是否要在第一阶段暴露采集状态 API。推荐后端先记录,UI 后续补齐。
|
||||
@@ -0,0 +1,197 @@
|
||||
# 规则流量清零设计
|
||||
|
||||
## 背景
|
||||
|
||||
Issue #523 希望“规则”页面中每条隧道规则显示的流量使用量支持手动清零。
|
||||
|
||||
当前规则流量保存在 `forward.in_flow` 和 `forward.out_flow`。流量上报时,同一份增量还会累计到用户总流量、用户隧道流量和相关配额统计中。因此,本功能必须将“规则展示计数器清零”与“用户或隧道配额重置”严格区分。
|
||||
|
||||
## 目标
|
||||
|
||||
为单条规则提供手动流量清零能力:
|
||||
|
||||
- 将所选规则的上传流量和下载流量清零。
|
||||
- 管理员可以清零任意规则。
|
||||
- 普通用户只能清零自己的规则。
|
||||
- 清零后,新产生的流量继续从零正常累计。
|
||||
|
||||
## 非目标
|
||||
|
||||
本功能不会:
|
||||
|
||||
- 修改用户总流量 `user.in_flow` 或 `user.out_flow`。
|
||||
- 修改用户隧道流量 `user_tunnel.in_flow` 或 `user_tunnel.out_flow`。
|
||||
- 修改每日或每月配额用量。
|
||||
- 修改历史流量统计。
|
||||
- 重置 nftables 节点计数器或其增量计算基线。
|
||||
- 重启、暂停、恢复或重新部署规则服务。
|
||||
- 增加批量流量清零功能。
|
||||
|
||||
## 后端设计
|
||||
|
||||
### API
|
||||
|
||||
新增接口:
|
||||
|
||||
```text
|
||||
POST /api/v1/forward/reset-flow
|
||||
```
|
||||
|
||||
请求体:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": 123
|
||||
}
|
||||
```
|
||||
|
||||
成功响应沿用统一 envelope:
|
||||
|
||||
```json
|
||||
{
|
||||
"code": 0,
|
||||
"msg": "success",
|
||||
"data": null,
|
||||
"ts": 0
|
||||
}
|
||||
```
|
||||
|
||||
具体 `msg`、`data` 和 `ts` 值继续由现有 response helper 生成。
|
||||
|
||||
### 参数与权限校验
|
||||
|
||||
Handler 执行以下步骤:
|
||||
|
||||
1. 只接受 `POST` 请求。
|
||||
2. 从 JSON 请求体读取正整数规则 ID。
|
||||
3. 调用现有 `resolveForwardAccess`:
|
||||
- 管理员角色可以访问任意存在的规则。
|
||||
- 普通用户仅能访问 `forward.user_id` 等于当前用户 ID 的规则。
|
||||
- 对普通用户访问他人规则的情况,沿用现有逻辑返回“转发不存在”,避免暴露规则存在性。
|
||||
4. 调用 Repository 完成清零。
|
||||
5. 返回统一成功响应。
|
||||
|
||||
### Repository
|
||||
|
||||
新增方法:
|
||||
|
||||
```go
|
||||
func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error
|
||||
```
|
||||
|
||||
该方法只更新指定 `forward` 记录:
|
||||
|
||||
```text
|
||||
in_flow = 0
|
||||
out_flow = 0
|
||||
updated_time = now
|
||||
```
|
||||
|
||||
Repository 不直接操作 Handler 的身份信息,也不更新任何其他表。
|
||||
|
||||
### 并发与后续流量
|
||||
|
||||
清零使用单条 SQL `UPDATE`。agent 流量上报和 nftables 流量采集仍使用原有增量累加逻辑。清零不会重置采集基线,因此下一次采集只会把清零之后新计算出的增量加回规则计数,不会把清零前的累计值整体恢复。
|
||||
|
||||
若清零 SQL 与流量增量 SQL 同时执行,数据库按实际语句执行顺序决定最终值;每条更新本身保持原子性。本功能不引入暂停采集或跨节点同步流程。
|
||||
|
||||
## 前端设计
|
||||
|
||||
### API 封装
|
||||
|
||||
在 `vite-frontend/src/api/index.ts` 新增:
|
||||
|
||||
```ts
|
||||
export const resetForwardFlow = (id: number) =>
|
||||
Network.post("/forward/reset-flow", { id });
|
||||
```
|
||||
|
||||
### 入口
|
||||
|
||||
在规则页面所有单条规则操作入口中增加“流量清零”操作:
|
||||
|
||||
- 分组表格视图。
|
||||
- 精简表格视图。
|
||||
- 卡片视图。
|
||||
|
||||
按钮使用独立的清零/刷新语义图标和提示文本,不复用删除按钮样式。
|
||||
|
||||
当规则的 `inFlow + outFlow` 等于零时,按钮禁用,避免重复请求。
|
||||
|
||||
### 确认交互
|
||||
|
||||
点击按钮后打开确认弹窗,显示规则名称,并明确说明:
|
||||
|
||||
- 仅清零当前规则显示的上传和下载流量。
|
||||
- 不影响用户总流量、用户隧道配额和历史统计。
|
||||
- 操作不可撤销。
|
||||
|
||||
确认期间显示 loading 状态并阻止重复提交。
|
||||
|
||||
### 成功与失败
|
||||
|
||||
- 成功:关闭弹窗,显示成功 toast,并刷新规则列表。
|
||||
- 失败:保留弹窗,显示后端错误信息或通用失败 toast。
|
||||
- 刷新后,该规则上传和下载均显示为零;后续流量继续正常累计。
|
||||
|
||||
## 错误处理
|
||||
|
||||
- 非 POST 请求:返回现有通用请求失败响应。
|
||||
- 请求体无法解析、ID 缺失或 ID 非正数:返回“请求参数错误”。
|
||||
- 规则不存在或普通用户访问他人规则:返回“转发不存在”。
|
||||
- Repository 更新失败:返回包含 Repository 错误信息的统一错误响应。
|
||||
- 前端网络错误:显示“流量清零失败”。
|
||||
|
||||
## 测试策略
|
||||
|
||||
### Repository 测试
|
||||
|
||||
验证:
|
||||
|
||||
- 指定规则的 `in_flow`、`out_flow` 被清零。
|
||||
- 指定规则的 `updated_time` 被更新。
|
||||
- 其他规则的流量不变。
|
||||
- 用户总流量不变。
|
||||
- 用户隧道流量不变。
|
||||
- Repository 未初始化时返回错误。
|
||||
|
||||
### Handler 测试
|
||||
|
||||
验证:
|
||||
|
||||
- 管理员能够清零任意存在的规则。
|
||||
- 普通用户能够清零自己的规则。
|
||||
- 普通用户不能清零他人的规则。
|
||||
- 不存在的规则返回错误。
|
||||
- 无效 ID 返回参数错误。
|
||||
- 非 POST 请求返回请求失败。
|
||||
- 成功请求不修改用户和用户隧道流量。
|
||||
|
||||
### 前端验证
|
||||
|
||||
项目没有配置前端测试框架,因此不新增前端单元测试。使用以下命令验证:
|
||||
|
||||
```bash
|
||||
(cd vite-frontend && pnpm run build)
|
||||
(cd vite-frontend && pnpm run lint)
|
||||
```
|
||||
|
||||
后端使用:
|
||||
|
||||
```bash
|
||||
(cd go-backend && go test ./...)
|
||||
```
|
||||
|
||||
## 文件范围
|
||||
|
||||
预计修改:
|
||||
|
||||
- `go-backend/internal/http/handler/handler.go`
|
||||
- `go-backend/internal/http/handler/mutations.go`
|
||||
- `go-backend/internal/http/handler/*_test.go`
|
||||
- `go-backend/internal/store/repo/repository_mutations.go`
|
||||
- `go-backend/internal/store/repo/*_test.go`
|
||||
- `vite-frontend/src/api/index.ts`
|
||||
- `vite-frontend/src/pages/forward.tsx`
|
||||
|
||||
不需要数据库迁移或新增依赖。
|
||||
@@ -7,7 +7,8 @@ RUN go mod download
|
||||
COPY . .
|
||||
ARG TARGETOS
|
||||
ARG TARGETARCH
|
||||
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -o /out/paneld ./cmd/paneld
|
||||
ARG KEYGEN_ACCOUNT_ID
|
||||
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -ldflags="-X 'go-backend/internal/license.AccountID=${KEYGEN_ACCOUNT_ID}'" -o /out/paneld ./cmd/paneld
|
||||
|
||||
FROM docker:27-cli AS dockercli
|
||||
|
||||
|
||||
@@ -29,6 +29,10 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
if value, gated := h.unlicensedPublicBrandValue(configName); gated {
|
||||
response.WriteJSON(w, response.OK(map[string]string{"name": configName, "value": value}))
|
||||
return
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(configName)
|
||||
if err != nil {
|
||||
@@ -42,3 +46,22 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
|
||||
|
||||
response.WriteJSON(w, response.OK(cfg))
|
||||
}
|
||||
|
||||
func (h *Handler) unlicensedPublicBrandValue(configName string) (string, bool) {
|
||||
switch configName {
|
||||
case "app_name", "app_logo", "app_favicon", "hide_footer_brand":
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
if isCommercial == "true" {
|
||||
return "", false
|
||||
}
|
||||
if configName == "app_name" {
|
||||
return "FLVX", true
|
||||
}
|
||||
if configName == "hide_footer_brand" {
|
||||
return "false", true
|
||||
}
|
||||
return "", true
|
||||
}
|
||||
|
||||
@@ -19,6 +19,8 @@ func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
|
||||
seedConfigValue(t, r, "app_logo", "logo-data")
|
||||
seedConfigValue(t, r, "app_favicon", "favicon-data")
|
||||
seedConfigValue(t, r, "app_bg_image", "bg-data")
|
||||
seedConfigValue(t, r, "app_bg_image_light", "light-bg-data")
|
||||
seedConfigValue(t, r, "app_bg_image_dark", "dark-bg-data")
|
||||
seedConfigValue(t, r, "cloudflare_site_key", "site-key")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
@@ -28,6 +30,50 @@ func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
for name, want := range map[string]string{
|
||||
"app_bg_image_light": "light-bg-data",
|
||||
"app_bg_image_dark": "dark-bg-data",
|
||||
} {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerConfigValue(t, resp, name, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicBrandConfigFallsBackWithoutCommercialLicense(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "app_name", "Paid Brand")
|
||||
seedConfigValue(t, r, "app_logo", "logo-data")
|
||||
seedConfigValue(t, r, "app_favicon", "favicon-data")
|
||||
seedConfigValue(t, r, "hide_footer_brand", "true")
|
||||
seedConfigValue(t, r, "is_commercial", "false")
|
||||
|
||||
for name, want := range map[string]string{
|
||||
"app_name": "FLVX",
|
||||
"app_logo": "",
|
||||
"app_favicon": "",
|
||||
"hide_footer_brand": "false",
|
||||
} {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerConfigValue(t, resp, name, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicBrandConfigUsesSavedValuesWithCommercialLicense(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
seedConfigValue(t, r, "app_name", "Paid Brand")
|
||||
seedConfigValue(t, r, "is_commercial", "true")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerConfigValue(t, resp, "app_name", "Paid Brand")
|
||||
}
|
||||
|
||||
func TestPublicConfigGetRejectsSensitiveKeys(t *testing.T) {
|
||||
@@ -82,6 +128,56 @@ func TestConfigGetAllowsSensitiveKeysForAdmin(t *testing.T) {
|
||||
assertHandlerConfigValue(t, resp, "jwt_secret", "jwt-secret")
|
||||
}
|
||||
|
||||
func TestConfigGetNeverReturnsLicenseCredentials(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
seedConfigValue(t, r, "license_key", "license-secret")
|
||||
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
|
||||
|
||||
for _, name := range []string{"license_key", "license_machine_id", "machine_fingerprint"} {
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
assertHandlerCodeMsg(t, resp, 403, "禁止访问系统授权凭据")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListNeverReturnsLicenseCredentials(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
seedConfigValue(t, r, "license_key", "license-secret")
|
||||
seedConfigValue(t, r, "license_machine_id", "machine-id")
|
||||
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
|
||||
seedConfigValue(t, r, "is_commercial", "true")
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/list", nil)
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
resp := httptest.NewRecorder()
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
var out struct {
|
||||
Code int `json:"code"`
|
||||
Data map[string]string `json:"data"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 || out.Data["is_commercial"] != "true" {
|
||||
t.Fatalf("unexpected config response: %+v", out)
|
||||
}
|
||||
if _, ok := out.Data["license_key"]; ok {
|
||||
t.Fatal("license_key must not be returned")
|
||||
}
|
||||
if _, ok := out.Data["license_machine_id"]; ok {
|
||||
t.Fatal("license_machine_id must not be returned")
|
||||
}
|
||||
if _, ok := out.Data["machine_fingerprint"]; ok {
|
||||
t.Fatal("machine_fingerprint must not be returned")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateAllowsSensitiveKeysForAdmin(t *testing.T) {
|
||||
router, _ := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
@@ -154,7 +250,7 @@ func TestConfigUpdateSingleAllowsCloudflareSecretKeyWrite(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateAllowsLicenseKeyWrite(t *testing.T) {
|
||||
func TestConfigUpdateRejectsLicenseKeyWrite(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
@@ -165,18 +261,13 @@ func TestConfigUpdateAllowsLicenseKeyWrite(t *testing.T) {
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
|
||||
cfg, err := r.GetConfigByName("license_key")
|
||||
if err != nil {
|
||||
t.Fatalf("get config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "license-secret" {
|
||||
t.Fatalf("expected license_key to be updated, got %#v", cfg)
|
||||
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
|
||||
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
|
||||
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigUpdateSingleAllowsLicenseKeyWrite(t *testing.T) {
|
||||
func TestConfigUpdateSingleRejectsLicenseKeyWrite(t *testing.T) {
|
||||
router, r := setupConfigAccessTestRouter(t)
|
||||
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
|
||||
|
||||
@@ -187,14 +278,9 @@ func TestConfigUpdateSingleAllowsLicenseKeyWrite(t *testing.T) {
|
||||
|
||||
router.ServeHTTP(resp, req)
|
||||
|
||||
assertHandlerCode(t, resp, 0)
|
||||
|
||||
cfg, err := r.GetConfigByName("license_key")
|
||||
if err != nil {
|
||||
t.Fatalf("get config: %v", err)
|
||||
}
|
||||
if cfg == nil || cfg.Value != "license-secret" {
|
||||
t.Fatalf("expected license_key to be updated, got %#v", cfg)
|
||||
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
|
||||
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
|
||||
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/client"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/ws"
|
||||
)
|
||||
@@ -58,6 +59,7 @@ type diagnosisWorkItem struct {
|
||||
type diagnosisExecOptions struct {
|
||||
commandTimeout time.Duration
|
||||
pingTimeoutMS int
|
||||
pingCount int
|
||||
timeoutMessage string
|
||||
}
|
||||
|
||||
@@ -772,6 +774,9 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
if payload, handled, err := h.diagnoseNftablesForwardRuntime(forward); handled || err != nil {
|
||||
return payload, err
|
||||
}
|
||||
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -787,6 +792,111 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (h *Handler) diagnoseNftablesForwardRuntime(forward *forwardRecord) (map[string]interface{}, bool, error) {
|
||||
if forward == nil {
|
||||
return nil, false, errForwardNotFound
|
||||
}
|
||||
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
if !nftMode {
|
||||
return nil, false, nil
|
||||
}
|
||||
if len(entryNodeIDs) == 0 {
|
||||
return nil, true, errors.New("nftables 转发缺少入口节点")
|
||||
}
|
||||
targets, err := resolveDiagnosisTargets(forward.RemoteAddr)
|
||||
if err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
results, err := h.buildNftablesForwardDiagnosisResults(forward, entryNodeIDs[0], targets)
|
||||
if err != nil {
|
||||
return nil, true, err
|
||||
}
|
||||
payload := map[string]interface{}{
|
||||
"forwardName": forward.Name,
|
||||
"timestamp": time.Now().UnixMilli(),
|
||||
"results": results,
|
||||
}
|
||||
return payload, true, nil
|
||||
}
|
||||
|
||||
func (h *Handler) buildNftablesForwardDiagnosisResults(forward *forwardRecord, nodeID int64, targets []diagnosisTarget) ([]map[string]interface{}, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return nil, errors.New("handler not initialized")
|
||||
}
|
||||
node, err := h.getNodeRecord(nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var binding *model.NftRuleBinding
|
||||
for i := range bindings {
|
||||
if bindings[i].ForwardID == forward.ID {
|
||||
binding = &bindings[i]
|
||||
break
|
||||
}
|
||||
}
|
||||
target := diagnosisTarget{}
|
||||
if len(targets) > 0 {
|
||||
target = targets[0]
|
||||
}
|
||||
|
||||
status := "missing"
|
||||
message := "nftables 规则未下发"
|
||||
success := false
|
||||
inPort := 0
|
||||
protocols := ""
|
||||
targetAddr := strings.TrimSpace(forward.RemoteAddr)
|
||||
ruleHash := ""
|
||||
if binding != nil {
|
||||
status = strings.ToLower(strings.TrimSpace(binding.Status))
|
||||
inPort = binding.InPort
|
||||
protocols = strings.TrimSpace(binding.Protocols)
|
||||
targetAddr = strings.TrimSpace(binding.TargetAddr)
|
||||
ruleHash = strings.TrimSpace(binding.RuleHash)
|
||||
if status == "" {
|
||||
status = "pending"
|
||||
}
|
||||
if status == runtimenft.StatusApplied {
|
||||
success = true
|
||||
message = "nftables 规则已下发"
|
||||
} else if strings.TrimSpace(binding.LastError) != "" {
|
||||
message = binding.LastError
|
||||
} else {
|
||||
message = "nftables 规则未完成下发"
|
||||
}
|
||||
}
|
||||
|
||||
packetLoss := 100
|
||||
if success {
|
||||
packetLoss = 0
|
||||
}
|
||||
result := map[string]interface{}{
|
||||
"success": success,
|
||||
"nodeName": node.Name,
|
||||
"nodeId": strconv.FormatInt(nodeID, 10),
|
||||
"targetIp": target.IP,
|
||||
"targetPort": target.Port,
|
||||
"description": fmt.Sprintf("nftables规则(%s)->目标(%s)", node.Name, defaultString(target.Address, targetAddr)),
|
||||
"averageTime": 0,
|
||||
"packetLoss": packetLoss,
|
||||
"message": message,
|
||||
"fromChainType": 1,
|
||||
"forwardMode": "nftables",
|
||||
"nftRuleStatus": status,
|
||||
"nftRuleHash": ruleHash,
|
||||
"inPort": inPort,
|
||||
"protocols": protocols,
|
||||
"targetAddr": targetAddr,
|
||||
}
|
||||
return []map[string]interface{}{result}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []diagnosisWorkItem, error) {
|
||||
if forward == nil {
|
||||
return "", nil, errForwardNotFound
|
||||
@@ -1487,10 +1597,14 @@ func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int, options diag
|
||||
if options.pingTimeoutMS <= 0 {
|
||||
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
|
||||
}
|
||||
pingCount := options.pingCount
|
||||
if pingCount <= 0 {
|
||||
pingCount = 4
|
||||
}
|
||||
res, err := h.sendNodeCommandWithTimeout(nodeID, "TcpPing", map[string]interface{}{
|
||||
"ip": ip,
|
||||
"port": port,
|
||||
"count": 4,
|
||||
"count": pingCount,
|
||||
"timeout": options.pingTimeoutMS,
|
||||
}, options.commandTimeout, false, false)
|
||||
if err != nil {
|
||||
@@ -1517,12 +1631,16 @@ func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int, op
|
||||
if options.pingTimeoutMS <= 0 {
|
||||
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
|
||||
}
|
||||
pingCount := options.pingCount
|
||||
if pingCount <= 0 {
|
||||
pingCount = 4
|
||||
}
|
||||
|
||||
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
|
||||
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
|
||||
IP: strings.TrimSpace(ip),
|
||||
Port: port,
|
||||
Count: 4,
|
||||
Count: pingCount,
|
||||
Timeout: options.pingTimeoutMS,
|
||||
Protocol: "tcp",
|
||||
})
|
||||
@@ -1757,6 +1875,7 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
services := make([]map[string]interface{}, 0, 2)
|
||||
targets := splitRemoteTargets(forward.RemoteAddr)
|
||||
strategy := strings.TrimSpace(forward.Strategy)
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(forward.ProxyProtocol, forward.ProxyProtocolReceive, forward.ProxyProtocolSend)
|
||||
if strategy == "" {
|
||||
strategy = "fifo"
|
||||
}
|
||||
@@ -1801,12 +1920,16 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
if runtimeLimiters.TrafficLimiter != "" {
|
||||
service["limiter"] = runtimeLimiters.TrafficLimiter
|
||||
}
|
||||
if forward.ProxyProtocol > 0 {
|
||||
if proxyProtocolReceive > 0 {
|
||||
serviceMetadata := ensureServiceMetadata(service)
|
||||
serviceMetadata["proxyProtocol"] = proxyProtocolReceive
|
||||
}
|
||||
if proxyProtocolSend > 0 {
|
||||
handlerConfig := service["handler"].(map[string]interface{})
|
||||
if handlerConfig["metadata"] == nil {
|
||||
handlerConfig["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
|
||||
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = proxyProtocolSend
|
||||
}
|
||||
if protocol == "udp" {
|
||||
listenerMetadata := map[string]interface{}{
|
||||
@@ -1819,10 +1942,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
|
||||
}
|
||||
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
|
||||
serviceMetadata := ensureServiceMetadata(service)
|
||||
serviceMetadata["interface"] = node.InterfaceName
|
||||
}
|
||||
services = append(services, service)
|
||||
}
|
||||
@@ -1841,6 +1962,25 @@ func buildForwarderNodes(targets []string) []map[string]interface{} {
|
||||
return nodes
|
||||
}
|
||||
|
||||
func ensureServiceMetadata(service map[string]interface{}) map[string]interface{} {
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
metadata, ok := service["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
metadata = map[string]interface{}{}
|
||||
service["metadata"] = metadata
|
||||
}
|
||||
return metadata
|
||||
}
|
||||
|
||||
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
|
||||
if send == 0 && legacy > 0 {
|
||||
send = legacy
|
||||
}
|
||||
return receive, send
|
||||
}
|
||||
|
||||
func processServerAddress(serverAddr string) string {
|
||||
serverAddr = normalizeServerAddressInput(serverAddr)
|
||||
if serverAddr == "" {
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -12,12 +13,27 @@ import (
|
||||
)
|
||||
|
||||
const bytesPerGB int64 = 1024 * 1024 * 1024
|
||||
const bytesPerMiB int64 = 1024 * 1024
|
||||
|
||||
func flowLimitBytes(flowGB, flowMiB int64) int64 {
|
||||
if flowMiB > 0 {
|
||||
if flowMiB > math.MaxInt64/bytesPerMiB {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return flowMiB * bytesPerMiB
|
||||
}
|
||||
if flowGB > math.MaxInt64/bytesPerGB {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return flowGB * bytesPerGB
|
||||
}
|
||||
|
||||
type userTunnelPolicy struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
TunnelID int64
|
||||
Flow int64
|
||||
FlowMiB int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
ExpTime int64
|
||||
@@ -358,7 +374,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
return errors.New("账号已过期")
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return errors.New("流量已超额,禁止开启转发")
|
||||
@@ -400,7 +416,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
return errors.New("该隧道已过期")
|
||||
}
|
||||
|
||||
utFlowLimit := policy.Flow * bytesPerGB
|
||||
utFlowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
|
||||
utCurrent := policy.InFlow + policy.OutFlow
|
||||
if utCurrent >= utFlowLimit {
|
||||
return errors.New("该隧道流量已超额,禁止开启转发")
|
||||
@@ -425,7 +441,7 @@ func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return true
|
||||
@@ -441,7 +457,7 @@ func shouldPauseUserTunnel(policy *userTunnelPolicy, now int64) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := policy.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
|
||||
current := policy.InFlow + policy.OutFlow
|
||||
if current >= flowLimit {
|
||||
return true
|
||||
@@ -465,7 +481,7 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er
|
||||
}
|
||||
return &userTunnelPolicy{
|
||||
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
|
||||
Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num,
|
||||
}, nil
|
||||
}
|
||||
@@ -659,9 +675,21 @@ func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName str
|
||||
}
|
||||
|
||||
func (h *Handler) speedLimiterExists(name string) bool {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
const forwardRulePrefix = "rule_traffic_limit_"
|
||||
if strings.HasPrefix(name, forwardRulePrefix) {
|
||||
forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64)
|
||||
if err != nil || forwardID <= 0 {
|
||||
return false
|
||||
}
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
return err == nil && forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0
|
||||
}
|
||||
|
||||
id, err := strconv.ParseInt(name, 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
return false
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestSpeedLimiterExistsPreservesForwardRuleLimiter(t *testing.T) {
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
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, ip_speed_id)
|
||||
VALUES(8, 1, 'user', 'forward', 1, '127.0.0.1:80', 'fifo', 0, 0, 1, 1, 1, 0, 3),
|
||||
(9, 1, 'user', 'forward-without-ip-limit', 1, '127.0.0.1:81', 'fifo', 0, 0, 1, 1, 1, 0, NULL)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert forward: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
if !h.speedLimiterExists("rule_traffic_limit_8") {
|
||||
t.Fatal("expected runtime limiter for existing forward to be preserved")
|
||||
}
|
||||
if h.speedLimiterExists("rule_traffic_limit_9") {
|
||||
t.Fatal("expected runtime limiter for forward without per-IP speed limit to be treated as orphaned")
|
||||
}
|
||||
if h.speedLimiterExists("rule_traffic_limit_10") {
|
||||
t.Fatal("expected runtime limiter for missing forward to be treated as orphaned")
|
||||
}
|
||||
if h.speedLimiterExists("rule_traffic_limit_invalid") {
|
||||
t.Fatal("expected malformed runtime limiter name to be treated as orphaned")
|
||||
}
|
||||
}
|
||||
@@ -9,14 +9,15 @@ import (
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing.T) {
|
||||
func TestBuildForwardServiceConfigsAppliesProxyProtocolReceiveAndSendIndependently(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocol: 2,
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocolReceive: 1,
|
||||
ProxyProtocolSend: 2,
|
||||
}
|
||||
tunnel := &tunnelRecord{Type: 1}
|
||||
node := &nodeRecord{
|
||||
@@ -38,8 +39,8 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
|
||||
if serviceMetadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface metadata eth0, got %v", serviceMetadata["interface"])
|
||||
}
|
||||
if _, ok := serviceMetadata["proxyProtocol"]; ok {
|
||||
t.Fatalf("proxyProtocol should not be listener metadata: %v", serviceMetadata)
|
||||
if serviceMetadata["proxyProtocol"] != 1 {
|
||||
t.Fatalf("expected service proxyProtocol 1 for receive mode, got %v", serviceMetadata["proxyProtocol"])
|
||||
}
|
||||
|
||||
handlerConfig, ok := service["handler"].(map[string]interface{})
|
||||
@@ -56,6 +57,46 @@ func TestBuildForwardServiceConfigsSendsProxyProtocolToForwardHandler(t *testing
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigsKeepsLegacyProxyProtocolAsSend(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocol: 2,
|
||||
}
|
||||
tunnel := &tunnelRecord{Type: 1}
|
||||
node := &nodeRecord{
|
||||
TCPListenAddr: "0.0.0.0",
|
||||
UDPListenAddr: "0.0.0.0",
|
||||
}
|
||||
|
||||
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", forwardRuntimeLimiters{})
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
|
||||
for _, service := range services {
|
||||
serviceMetadata, _ := service["metadata"].(map[string]interface{})
|
||||
if _, ok := serviceMetadata["proxyProtocol"]; ok {
|
||||
t.Fatalf("legacy proxyProtocol should not enable receive mode: %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 legacy proxyProtocol to send version 2, got %v", handlerMetadata["proxyProtocol"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestForwardResetFlowPermissionsAndIsolation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
actorID int64
|
||||
actorRole int
|
||||
forwardID int64
|
||||
wantCode int
|
||||
wantInFlow int64
|
||||
wantOutFlow int64
|
||||
}{
|
||||
{name: "admin resets another user's rule", actorID: 1, actorRole: 0, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "owner resets own rule", actorID: 2, actorRole: 1, forwardID: 20, wantCode: 0, wantInFlow: 0, wantOutFlow: 0},
|
||||
{name: "user cannot reset another user's rule", actorID: 3, actorRole: 1, forwardID: 20, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
{name: "missing rule is rejected", actorID: 1, actorRole: 0, forwardID: 999, wantCode: -1, wantInFlow: 111, wantOutFlow: 222},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
h, r := setupForwardResetFlowHandler(t)
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, tt.forwardID, tt.actorID, tt.actorRole)
|
||||
res := httptest.NewRecorder()
|
||||
|
||||
h.forwardResetFlow(res, req)
|
||||
|
||||
if got := decodeForwardResetFlowCode(t, res); got != tt.wantCode {
|
||||
t.Fatalf("code = %d, want %d; body=%s", got, tt.wantCode, res.Body.String())
|
||||
}
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", tt.wantInFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", tt.wantOutFlow)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardResetFlowDBValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardResetFlowRejectsInvalidRequests(t *testing.T) {
|
||||
h, _ := setupForwardResetFlowHandler(t)
|
||||
|
||||
t.Run("non post", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodGet, 20, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("invalid id", func(t *testing.T) {
|
||||
req := newForwardResetFlowRequest(t, http.MethodPost, 0, 1, 0)
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardResetFlow(res, req)
|
||||
if code := decodeForwardResetFlowCode(t, res); code != -1 {
|
||||
t.Fatalf("code = %d, want -1", code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func setupForwardResetFlowHandler(t *testing.T) (*Handler, *repo.Repository) {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "forward-reset-handler.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
statements := []string{
|
||||
`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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)`,
|
||||
`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, 'other', 'pwd', 1, 0, 100, 0, 0, 0, 10, 1000, 1000, 1)`,
|
||||
`INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(1, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)`,
|
||||
`INSERT INTO user_tunnel(id, user_id, tunnel_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, 10, 100, 500, 600, 0, 0, 1)`,
|
||||
`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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, 1000, 1, 0)`,
|
||||
}
|
||||
for _, statement := range statements {
|
||||
if err := r.DB().Exec(statement).Error; err != nil {
|
||||
t.Fatalf("seed database: %v", err)
|
||||
}
|
||||
}
|
||||
return New(r, "test-secret"), r
|
||||
}
|
||||
|
||||
func newForwardResetFlowRequest(t *testing.T, method string, forwardID, actorID int64, roleID int) *http.Request {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(map[string]int64{"id": forwardID})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal request: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(method, "/api/v1/forward/reset-flow", bytes.NewReader(body))
|
||||
claims := auth.Claims{Sub: strconv.FormatInt(actorID, 10), RoleID: roleID}
|
||||
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
|
||||
}
|
||||
|
||||
func decodeForwardResetFlowCode(t *testing.T, res *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
}
|
||||
if err := json.Unmarshal(res.Body.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("decode response: %v; body=%s", err, res.Body.String())
|
||||
}
|
||||
return payload.Code
|
||||
}
|
||||
|
||||
func assertForwardResetFlowDBValue(t *testing.T, r *repo.Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
@@ -20,7 +21,6 @@ import (
|
||||
"go-backend/internal/health"
|
||||
"go-backend/internal/http/middleware"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/monitoring"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
@@ -42,10 +42,12 @@ type Handler struct {
|
||||
captchaMu sync.Mutex
|
||||
captchaTokens map[string]int64
|
||||
|
||||
jobsMu sync.Mutex
|
||||
jobsCancel context.CancelFunc
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
jobsMu sync.Mutex
|
||||
jobsCancel context.CancelFunc
|
||||
jobsStarted bool
|
||||
jobsWG sync.WaitGroup
|
||||
fingerprintMu sync.Mutex
|
||||
licenseValidationMu sync.Mutex
|
||||
|
||||
upgradeMu sync.Mutex
|
||||
systemUpgradeMu sync.Mutex
|
||||
@@ -221,6 +223,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/forward/force-delete", h.forwardForceDelete)
|
||||
mux.HandleFunc("/api/v1/forward/pause", h.forwardPause)
|
||||
mux.HandleFunc("/api/v1/forward/resume", h.forwardResume)
|
||||
mux.HandleFunc("/api/v1/forward/reset-flow", h.forwardResetFlow)
|
||||
mux.HandleFunc("/api/v1/forward/diagnose", h.forwardDiagnose)
|
||||
mux.HandleFunc("/api/v1/forward/diagnose/stream", h.forwardDiagnoseStream)
|
||||
mux.HandleFunc("/api/v1/forward/update-order", h.forwardUpdateOrder)
|
||||
@@ -399,6 +402,10 @@ func (h *Handler) getConfigByName(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
configName := strings.ToLower(strings.TrimSpace(req.Name))
|
||||
if configName == "license_key" || configName == "license_machine_id" || configName == "machine_fingerprint" {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问系统授权凭据"))
|
||||
return
|
||||
}
|
||||
if repo.IsSensitiveConfigKey(configName) && !isAdminRequest(r) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
@@ -434,11 +441,13 @@ func (h *Handler) getConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
ctxClaims := r.Context().Value(middleware.ClaimsContextKey)
|
||||
if claims, ok := ctxClaims.(auth.Claims); !ok || claims.RoleID != 0 {
|
||||
delete(cfgMap, "license_key")
|
||||
delete(cfgMap, "cloudflare_secret_key")
|
||||
delete(cfgMap, "jwt_secret")
|
||||
claims, isAdmin := ctxClaims.(auth.Claims)
|
||||
if !isAdmin || claims.RoleID != 0 {
|
||||
cfgMap = repo.FilterSensitiveConfigs(cfgMap)
|
||||
}
|
||||
delete(cfgMap, "license_key")
|
||||
delete(cfgMap, "license_machine_id")
|
||||
delete(cfgMap, "machine_fingerprint")
|
||||
response.WriteJSON(w, response.OK(cfgMap))
|
||||
}
|
||||
|
||||
@@ -710,6 +719,7 @@ func (h *Handler) userTunnelList(w http.ResponseWriter, r *http.Request) {
|
||||
"tunnelName": t.TunnelName,
|
||||
"status": t.Status,
|
||||
"flow": t.Flow,
|
||||
"flowMiB": t.FlowMiB,
|
||||
"num": t.Num,
|
||||
"expTime": t.ExpTime,
|
||||
"flowResetTime": t.FlowResetTime,
|
||||
@@ -877,10 +887,16 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
func (h *Handler) getOrCreateMachineFingerprint() (string, error) {
|
||||
fp, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
h.fingerprintMu.Lock()
|
||||
defer h.fingerprintMu.Unlock()
|
||||
|
||||
fp, err := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
if fp != "" {
|
||||
return fp, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return "", err
|
||||
}
|
||||
|
||||
newFp := uuid.New().String()
|
||||
now := time.Now().UnixMilli()
|
||||
@@ -907,56 +923,35 @@ func (h *Handler) licenseActivate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("授权码不能为空"))
|
||||
return
|
||||
}
|
||||
h.licenseValidationMu.Lock()
|
||||
defer h.licenseValidationMu.Unlock()
|
||||
|
||||
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
|
||||
|
||||
fingerprint, err := h.getOrCreateMachineFingerprint()
|
||||
valResp, err := h.validateLicenseForMachine(key)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("生成设备指纹失败"))
|
||||
return
|
||||
}
|
||||
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("连接授权服务器失败: "+err.Error()))
|
||||
log.Printf("license activation failed: %v", err)
|
||||
response.WriteJSON(w, response.ErrDefault(licenseValidationErrorMessage(err)))
|
||||
return
|
||||
}
|
||||
|
||||
if !valResp.Meta.Valid {
|
||||
if valResp.Meta.Code == "NO_MACHINES" || valResp.Meta.Code == "NO_MACHINE" || valResp.Meta.Code == "MACHINE_SCOPE_REQUIRED" || valResp.Meta.Code == "FINGERPRINT_SCOPE_MISMATCH" {
|
||||
// Needs machine activation
|
||||
client.Token = key
|
||||
err = client.ActivateMachine(valResp.Data.ID, fingerprint)
|
||||
if err != nil {
|
||||
// Translate specific error messages or log them
|
||||
response.WriteJSON(w, response.ErrDefault("设备绑定失败: "+err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
// Validation might still fail with scope if we don't query via machine id, but since activate machine succeeded
|
||||
// we can consider the license valid for our simple usecase
|
||||
} else {
|
||||
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
|
||||
return
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.UpsertConfig("license_key", key, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpsertConfig("is_commercial", "true", now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
expiry := valResp.Data.Attributes.Expiry
|
||||
if expiry == "" {
|
||||
expiry = "never"
|
||||
}
|
||||
if err := h.repo.UpsertConfig("license_expiry", expiry, now); err != nil {
|
||||
licenseState := map[string]string{
|
||||
"license_key": key,
|
||||
"is_commercial": "true",
|
||||
"license_expiry": expiry,
|
||||
}
|
||||
if valResp.MachineID != "" {
|
||||
licenseState["license_machine_id"] = valResp.MachineID
|
||||
}
|
||||
if err := h.repo.UpsertConfigs(licenseState, now); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -998,6 +993,10 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
if repo.IsSystemManagedConfigKey(key) {
|
||||
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
|
||||
return
|
||||
}
|
||||
|
||||
if protectedKeys[key] && isCommercial != "true" {
|
||||
response.WriteJSON(w, response.ErrDefault("需要商业版授权"))
|
||||
@@ -1014,6 +1013,7 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.notifyTunnelQualityConfigChanged(key)
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
@@ -1039,6 +1039,10 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||
return
|
||||
}
|
||||
if repo.IsSystemManagedConfigKey(name) {
|
||||
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
|
||||
return
|
||||
}
|
||||
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
if (name == "app_name" || name == "app_logo" || name == "app_favicon" || name == "hide_footer_brand") && isCommercial != "true" {
|
||||
@@ -1061,6 +1065,7 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
h.notifyTunnelQualityConfigChanged(name)
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
@@ -1109,11 +1114,23 @@ func normalizeAndValidateConfigValue(key, value string) (string, error) {
|
||||
}
|
||||
case monitoring.ConfigMonitorRetentionDays:
|
||||
return monitoring.NormalizeMonitoringRetentionDays(value)
|
||||
case monitoring.ConfigTunnelQualityProbeIntervalSec:
|
||||
return monitoring.NormalizeTunnelQualityProbeIntervalSeconds(value)
|
||||
default:
|
||||
return value, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) notifyTunnelQualityConfigChanged(key string) {
|
||||
if h == nil || h.qualityProber == nil {
|
||||
return
|
||||
}
|
||||
switch strings.TrimSpace(key) {
|
||||
case monitorTunnelQualityEnabledConfigKey, monitoring.ConfigTunnelQualityProbeIntervalSec:
|
||||
h.qualityProber.NotifyConfigChanged()
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return true
|
||||
@@ -1197,6 +1214,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
"tunnelName": t.TunnelName,
|
||||
"tunnelFlow": t.TunnelFlow,
|
||||
"flow": t.Flow,
|
||||
"flowMiB": t.FlowMiB,
|
||||
"inFlow": t.InFlow,
|
||||
"outFlow": t.OutFlow,
|
||||
"num": t.Num,
|
||||
@@ -1246,6 +1264,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
"user": user.User,
|
||||
"status": user.Status,
|
||||
"flow": user.Flow,
|
||||
"flowMiB": user.FlowMiB,
|
||||
"inFlow": user.InFlow,
|
||||
"outFlow": user.OutFlow,
|
||||
"num": user.Num,
|
||||
|
||||
@@ -2,11 +2,12 @@ package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/license"
|
||||
)
|
||||
|
||||
var nftablesTrafficCollectInterval = 30 * time.Second
|
||||
|
||||
func (h *Handler) StartBackgroundJobs() {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
@@ -20,7 +21,7 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
h.jobsCancel = cancel
|
||||
h.jobsStarted = true
|
||||
h.jobsWG.Add(7)
|
||||
h.jobsWG.Add(8)
|
||||
h.jobsMu.Unlock()
|
||||
|
||||
go h.runHourlyStatsLoop(ctx)
|
||||
@@ -30,10 +31,12 @@ func (h *Handler) StartBackgroundJobs() {
|
||||
go h.runHealthChecks(ctx)
|
||||
go h.runTunnelQualityProber(ctx)
|
||||
go h.runValidateLicenseJob(ctx)
|
||||
go h.runNftablesTrafficCollectLoop(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runValidateLicenseJob(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
h.validateLicenseJob()
|
||||
ticker := time.NewTicker(12 * time.Hour)
|
||||
defer ticker.Stop()
|
||||
|
||||
@@ -51,22 +54,25 @@ func (h *Handler) validateLicenseJob() {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
|
||||
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
|
||||
h.licenseValidationMu.Lock()
|
||||
defer h.licenseValidationMu.Unlock()
|
||||
|
||||
key, _ := h.repo.GetViteConfigValue("license_key")
|
||||
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
|
||||
|
||||
if key == "" || isCommercial != "true" {
|
||||
if key == "" {
|
||||
return // Nothing to validate
|
||||
}
|
||||
|
||||
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
|
||||
client := license.NewKeygenClient(accountID, "")
|
||||
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
|
||||
valResp, err := h.validateLicenseForMachine(key)
|
||||
|
||||
if err != nil {
|
||||
// Network error or timeout. Grace period by not revoking immediately here.
|
||||
// Network and decode failures have no validation response, so retain the
|
||||
// current state as a grace period. A rejected machine binding still has
|
||||
// the original invalid response and must not stay commercially enabled.
|
||||
if licenseValidationErrorIsDefinitive(valResp, err) {
|
||||
now := time.Now().UnixMilli()
|
||||
_ = h.repo.UpsertConfig("is_commercial", "false", now)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
@@ -80,7 +86,14 @@ func (h *Handler) validateLicenseJob() {
|
||||
if expiry == "" {
|
||||
expiry = "never"
|
||||
}
|
||||
_ = h.repo.UpsertConfig("license_expiry", expiry, now)
|
||||
licenseState := map[string]string{
|
||||
"is_commercial": "true",
|
||||
"license_expiry": expiry,
|
||||
}
|
||||
if valResp.MachineID != "" {
|
||||
licenseState["license_machine_id"] = valResp.MachineID
|
||||
}
|
||||
_ = h.repo.UpsertConfigs(licenseState, now)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -128,6 +141,54 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) {
|
||||
h.qualityProber.Start(ctx)
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
h.runNftablesStartupReconcile(ctx)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
h.runNftablesTrafficCollectJob(time.Now())
|
||||
}
|
||||
|
||||
interval := nftablesTrafficCollectInterval
|
||||
if interval <= 0 {
|
||||
interval = 30 * time.Second
|
||||
}
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
h.runNftablesTrafficCollectJob(time.Now())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesStartupReconcile(ctx context.Context) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
nodes, err := h.repo.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
log.Printf("nftables startup reconcile failed op=list_nodes err=%v", err)
|
||||
return
|
||||
}
|
||||
for _, node := range nodes {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
if err := h.syncNftablesNode(node.NodeID); err != nil {
|
||||
log.Printf("nftables startup reconcile failed node_id=%d err=%v", node.NodeID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
|
||||
defer h.jobsWG.Done()
|
||||
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"go-backend/internal/license"
|
||||
)
|
||||
|
||||
var newLicenseClient = license.NewKeygenClient
|
||||
|
||||
func keygenAccountID() string {
|
||||
if value := strings.TrimSpace(license.AccountID); value != "" {
|
||||
return value
|
||||
}
|
||||
return strings.TrimSpace(os.Getenv("KEYGEN_ACCOUNT_ID"))
|
||||
}
|
||||
|
||||
func licenseNeedsMachineActivation(code string) bool {
|
||||
switch strings.ToUpper(strings.TrimSpace(code)) {
|
||||
case "NO_MACHINES", "NO_MACHINE", "MACHINE_SCOPE_REQUIRED", "FINGERPRINT_SCOPE_MISMATCH":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func licenseValidationErrorIsDefinitive(validation *license.ValidateResponse, err error) bool {
|
||||
if err == nil {
|
||||
return validation != nil && !validation.Meta.Valid
|
||||
}
|
||||
var apiErr *license.APIError
|
||||
if !errors.As(err, &apiErr) {
|
||||
return false
|
||||
}
|
||||
if apiErr.StatusCode == http.StatusTooManyRequests || apiErr.StatusCode >= http.StatusInternalServerError {
|
||||
return false
|
||||
}
|
||||
return apiErr.Operation == "activate machine" && validation != nil && !validation.Meta.Valid
|
||||
}
|
||||
|
||||
func licenseValidationErrorMessage(err error) string {
|
||||
if strings.Contains(err.Error(), "keygen account id is not configured") {
|
||||
return "授权服务配置错误"
|
||||
}
|
||||
var apiErr *license.APIError
|
||||
if errors.As(err, &apiErr) {
|
||||
if apiErr.HasCode("MACHINE_LIMIT_EXCEEDED") {
|
||||
return "授权设备数量已达上限"
|
||||
}
|
||||
if apiErr.StatusCode == http.StatusUnauthorized || apiErr.StatusCode == http.StatusForbidden {
|
||||
return "授权码无效或无权绑定设备"
|
||||
}
|
||||
}
|
||||
return "连接授权服务器失败,请稍后重试"
|
||||
}
|
||||
|
||||
func (h *Handler) validateLicenseForMachine(key string) (*license.ValidateResponse, error) {
|
||||
fingerprint, err := h.getOrCreateMachineFingerprint()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prepare machine fingerprint: %w", err)
|
||||
}
|
||||
storedMachineID, _ := h.repo.GetViteConfigValue("license_machine_id")
|
||||
|
||||
accountID := keygenAccountID()
|
||||
if accountID == "" {
|
||||
return nil, fmt.Errorf("keygen account id is not configured")
|
||||
}
|
||||
client := newLicenseClient(accountID, "")
|
||||
var validation *license.ValidateResponse
|
||||
if storedMachineID != "" {
|
||||
validation, err = client.ValidateKeyWithMachine(key, fingerprint, storedMachineID)
|
||||
} else {
|
||||
validation, err = client.ValidateKeyWithFingerprint(key, fingerprint)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validation.Meta.Valid || !licenseNeedsMachineActivation(validation.Meta.Code) {
|
||||
if validation.Meta.Valid {
|
||||
validation.MachineID = storedMachineID
|
||||
}
|
||||
return validation, nil
|
||||
}
|
||||
|
||||
client.Token = key
|
||||
machineID, err := client.ActivateMachine(validation.Data.ID, fingerprint)
|
||||
if err != nil {
|
||||
return validation, err
|
||||
}
|
||||
if machineID == "" {
|
||||
machineID, err = client.GetMachineID(fingerprint)
|
||||
if err != nil {
|
||||
return validation, fmt.Errorf("retrieve activated machine: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
validation, err = client.ValidateKeyWithMachine(key, fingerprint, machineID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validation.Meta.Valid {
|
||||
validation.MachineID = machineID
|
||||
}
|
||||
return validation, nil
|
||||
}
|
||||
@@ -0,0 +1,358 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestValidateLicenseJobRepairsMissingMachineBinding(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
|
||||
var validations atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
if validations.Add(1) == 1 {
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
return
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
|
||||
fingerprint, err := r.GetViteConfigValue("machine_fingerprint")
|
||||
if err != nil || strings.TrimSpace(fingerprint) == "" {
|
||||
t.Fatalf("expected persisted machine fingerprint, got value=%q err=%v", fingerprint, err)
|
||||
}
|
||||
if got := validations.Load(); got != 2 {
|
||||
t.Fatalf("validation calls = %d, want 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobAcceptsExistingMachineActivation(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
|
||||
var validations atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
if validations.Add(1) == 1 {
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
return
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "never")
|
||||
if got := validations.Load(); got != 2 {
|
||||
t.Fatalf("validation calls = %d, want 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivateRequiresSuccessfulPostActivationValidation(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
|
||||
var validations atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
code := "NO_MACHINE"
|
||||
if validations.Add(1) > 1 {
|
||||
code = "FINGERPRINT_SCOPE_MISMATCH"
|
||||
}
|
||||
_, _ = fmt.Fprintf(w, `{"meta":{"valid":false,"code":%q},"data":{"id":"license-id","attributes":{}}}`, code)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
|
||||
res := httptest.NewRecorder()
|
||||
h.licenseActivate(res, req)
|
||||
|
||||
if !strings.Contains(res.Body.String(), "FINGERPRINT_SCOPE_MISMATCH") {
|
||||
t.Fatalf("expected post-activation validation failure, got %s", res.Body.String())
|
||||
}
|
||||
assertLicenseConfig(t, r, "is_commercial", "false")
|
||||
for _, name := range []string{"license_key", "license_expiry"} {
|
||||
if value, err := r.GetViteConfigValue(name); err == nil || value != "" {
|
||||
t.Fatalf("%s should not be persisted, got value=%q err=%v", name, value, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivatePersistsValidatedState(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
|
||||
res := httptest.NewRecorder()
|
||||
h.licenseActivate(res, req)
|
||||
|
||||
if !strings.Contains(res.Body.String(), `"code":0`) {
|
||||
t.Fatalf("expected activation success, got %s", res.Body.String())
|
||||
}
|
||||
assertLicenseConfig(t, r, "license_key", "license-secret")
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobDowngradesWhenMachineBindingIsRejected(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "false")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobRestoresCommercialStateWhenLicenseRecovers(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "false", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
if strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, req)
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_expiry", "never")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobUsesStoredMachineScope(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
seedLicenseConfig(t, r, "license_machine_id", "machine-id", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
if !strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
|
||||
t.Fatalf("unexpected request %s %s", req.Method, req.URL.Path)
|
||||
}
|
||||
var body struct {
|
||||
Meta struct {
|
||||
Scope map[string]string `json:"scope"`
|
||||
} `json:"meta"`
|
||||
}
|
||||
if err := json.NewDecoder(req.Body).Decode(&body); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
if body.Meta.Scope["machine"] != "machine-id" || body.Meta.Scope["fingerprint"] != "fingerprint" {
|
||||
t.Fatalf("unexpected validation scope: %+v", body.Meta.Scope)
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
assertLicenseConfig(t, r, "license_machine_id", "machine-id")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobKeepsStateOnMachineLookupServerFailure(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"SERVICE_UNAVAILABLE"}]}`)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
}
|
||||
|
||||
func TestValidateLicenseJobKeepsStateOnMachineLookupNotFound(t *testing.T) {
|
||||
r := openLicenseTestRepository(t)
|
||||
now := time.Now().UnixMilli()
|
||||
seedLicenseConfig(t, r, "license_key", "license-secret", now)
|
||||
seedLicenseConfig(t, r, "is_commercial", "true", now)
|
||||
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
case strings.HasSuffix(req.URL.Path, "/machines"):
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
|
||||
case strings.Contains(req.URL.Path, "/machines/"):
|
||||
http.NotFound(w, req)
|
||||
default:
|
||||
http.NotFound(w, req)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
restoreLicenseClientFactory(t, server.URL)
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.validateLicenseJob()
|
||||
|
||||
assertLicenseConfig(t, r, "is_commercial", "true")
|
||||
}
|
||||
|
||||
func TestLicenseValidationErrorMessageDoesNotExposeKeygenResponse(t *testing.T) {
|
||||
err := &license.APIError{
|
||||
Operation: "activate machine",
|
||||
StatusCode: http.StatusUnprocessableEntity,
|
||||
Body: `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED","detail":"private detail"}]}`,
|
||||
}
|
||||
message := licenseValidationErrorMessage(err)
|
||||
if message != "授权设备数量已达上限" || strings.Contains(message, "private detail") {
|
||||
t.Fatalf("unexpected public error message %q", message)
|
||||
}
|
||||
}
|
||||
|
||||
func openLicenseTestRepository(t *testing.T) *repo.Repository {
|
||||
t.Helper()
|
||||
r, err := repo.Open(filepath.Join(t.TempDir(), "license.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("repo.Open() error = %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
return r
|
||||
}
|
||||
|
||||
func seedLicenseConfig(t *testing.T, r *repo.Repository, name, value string, now int64) {
|
||||
t.Helper()
|
||||
if err := r.UpsertConfig(name, value, now); err != nil {
|
||||
t.Fatalf("UpsertConfig(%q) error = %v", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertLicenseConfig(t *testing.T, r *repo.Repository, name, want string) {
|
||||
t.Helper()
|
||||
got, err := r.GetViteConfigValue(name)
|
||||
if err != nil {
|
||||
t.Fatalf("GetViteConfigValue(%q) error = %v", name, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("config %q = %q, want %q", name, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func restoreLicenseClientFactory(t *testing.T, baseURL string) {
|
||||
t.Helper()
|
||||
t.Setenv("KEYGEN_ACCOUNT_ID", "account-id")
|
||||
previous := newLicenseClient
|
||||
newLicenseClient = func(accountID, token string) *license.KeygenClient {
|
||||
client := license.NewKeygenClient(accountID, token)
|
||||
client.BaseURL = baseURL
|
||||
return client
|
||||
}
|
||||
t.Cleanup(func() { newLicenseClient = previous })
|
||||
}
|
||||
@@ -57,7 +57,11 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
status := asInt(req["status"], 1)
|
||||
flow := asInt64(req["flow"], 100)
|
||||
flow, flowMiB, flowErr := parseTrafficLimit(req, 100)
|
||||
if flowErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
|
||||
return
|
||||
}
|
||||
num := asInt(req["num"], 10)
|
||||
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
|
||||
flowResetTime := asInt64(req["flowResetTime"], 1)
|
||||
@@ -76,7 +80,7 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
|
||||
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now, flowMiB)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -164,7 +168,16 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
flow := asInt64(req["flow"], 100)
|
||||
flow, flowMiB, flowErr := parseTrafficLimit(req, 100)
|
||||
if flowErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
|
||||
return
|
||||
}
|
||||
if _, supplied := req["flowMiB"]; !supplied {
|
||||
if current, err := h.repo.GetUserByID(id); err == nil && current != nil && current.Flow == flow {
|
||||
flowMiB = current.FlowMiB
|
||||
}
|
||||
}
|
||||
num := asInt(req["num"], 10)
|
||||
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
|
||||
flowResetTime := asInt64(req["flowResetTime"], 1)
|
||||
@@ -176,7 +189,7 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
pwd := asString(req["pwd"])
|
||||
if strings.TrimSpace(pwd) == "" {
|
||||
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now, flowMiB); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -186,13 +199,13 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now, flowMiB); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime)
|
||||
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime, flowMiB)
|
||||
if hasDailyQuota || hasMonthlyQuota {
|
||||
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
|
||||
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
|
||||
@@ -417,6 +430,12 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
newHTTP := asInt(req["http"], currentHTTP)
|
||||
newTLS := asInt(req["tls"], currentTLS)
|
||||
newSocks := asInt(req["socks"], currentSocks)
|
||||
forwardMode := defaultNodeForwardMode(strings.TrimSpace(asString(req["forwardMode"])))
|
||||
currentForwardMode, err := h.repo.GetNodeForwardMode(id)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
serverIP := asString(req["serverIp"])
|
||||
if serverIP != "" {
|
||||
if err := IsValidNodeAddress(serverIP); err != nil {
|
||||
@@ -424,7 +443,8 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
}
|
||||
if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
|
||||
usesNftablesRuntime := forwardMode == "nftables" || defaultNodeForwardMode(currentForwardMode) == "nftables"
|
||||
if currentStatus == 1 && !usesNftablesRuntime && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
|
||||
if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
@@ -432,7 +452,6 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwardMode := defaultNodeForwardMode(strings.TrimSpace(asString(req["forwardMode"])))
|
||||
if err := h.repo.UpdateNode(id,
|
||||
asString(req["name"]),
|
||||
serverIP,
|
||||
@@ -1985,14 +2004,32 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, oldErr.Error()))
|
||||
return
|
||||
}
|
||||
oldTunnel, oldTunnelErr := h.repo.GetUserTunnelByID(id)
|
||||
if oldTunnelErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, oldTunnelErr.Error()))
|
||||
return
|
||||
}
|
||||
if oldTunnel == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道权限不存在"))
|
||||
return
|
||||
}
|
||||
flow, flowMiB, flowErr := parseTrafficLimit(req, 0)
|
||||
if flowErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
|
||||
return
|
||||
}
|
||||
if _, supplied := req["flowMiB"]; !supplied && oldTunnel.Flow == flow {
|
||||
flowMiB = oldTunnel.FlowMiB
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserTunnel(id,
|
||||
asInt64(req["flow"], 0),
|
||||
flow,
|
||||
asInt(req["num"], 0),
|
||||
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
|
||||
asInt64(req["flowResetTime"], 1),
|
||||
nullableInt(speedID),
|
||||
asInt(req["status"], 1),
|
||||
flowMiB,
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -2007,6 +2044,7 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
oldFlowReset,
|
||||
oldSpeedID,
|
||||
oldStatus,
|
||||
oldTunnel.FlowMiB,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr)))
|
||||
@@ -2143,8 +2181,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
ipMaxConn = 0
|
||||
}
|
||||
proxyProtocol := asInt(req["proxyProtocol"], 0)
|
||||
proxyProtocolReceive := asInt(req["proxyProtocolReceive"], 0)
|
||||
proxyProtocolSend := asInt(req["proxyProtocolSend"], 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)
|
||||
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, proxyProtocolReceive, proxyProtocolSend)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -2333,8 +2373,10 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
ipMaxConn = 0
|
||||
}
|
||||
proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol)
|
||||
proxyProtocolReceive := asInt(req["proxyProtocolReceive"], forward.ProxyProtocolReceive)
|
||||
proxyProtocolSend := asInt(req["proxyProtocolSend"], forward.ProxyProtocolSend)
|
||||
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol); err != nil {
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, ipMaxConn, newIPSpeedID, proxyProtocol, proxyProtocolReceive, proxyProtocolSend); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -2412,23 +2454,30 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
var nftNodeID int64
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, modeErr.Error()))
|
||||
return
|
||||
} else if nftMode && len(entryNodeIDs) > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("nftables 转发缺少入口节点"))
|
||||
return
|
||||
}
|
||||
nftNodeID = entryNodeIDs[0]
|
||||
} else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.deleteForwardByID(id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if nftNodeID > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2485,6 +2534,26 @@ func (h *Handler) forwardPause(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) forwardResetFlow(w http.ResponseWriter, r *http.Request) {
|
||||
id := idFromBody(r, w)
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
if _, _, _, err := h.resolveForwardAccess(r, id); err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.ResetForwardFlow(id, time.Now().UnixMilli()); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) {
|
||||
id := idFromBody(r, w)
|
||||
if id <= 0 {
|
||||
@@ -2577,7 +2646,19 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
var nftNodeID int64
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
f++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
nftNodeID = entryNodeIDs[0]
|
||||
} else if err := h.controlForwardServices(forward, "DeleteService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
@@ -2585,9 +2666,16 @@ func (h *Handler) forwardBatchDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if err := h.deleteForwardByID(id); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
continue
|
||||
}
|
||||
if nftNodeID > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(nftNodeID); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
continue
|
||||
}
|
||||
}
|
||||
s++
|
||||
}
|
||||
response.WriteJSON(w, response.OK(batchOperationResult{SuccessCount: s, FailCount: f, Failures: failures}))
|
||||
}
|
||||
@@ -4696,7 +4784,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.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.IPMaxConn, oldForward.IPSpeedID, oldForward.ProxyProtocol, oldForward.ProxyProtocolReceive, oldForward.ProxyProtocolSend,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
@@ -4725,6 +4813,14 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
}
|
||||
|
||||
reqFlow := asInt64(req["flow"], -1)
|
||||
var reqFlowMiB int64
|
||||
if _, hasFlowMiB := req["flowMiB"]; hasFlowMiB {
|
||||
var flowErr error
|
||||
reqFlow, reqFlowMiB, flowErr = parseTrafficLimit(req, 0)
|
||||
if flowErr != nil {
|
||||
return flowErr
|
||||
}
|
||||
}
|
||||
reqNum := asInt(req["num"], -1)
|
||||
reqExpTime := asInt64(req["expTime"], -1)
|
||||
reqFlowReset := asInt64(req["flowResetTime"], -1)
|
||||
@@ -4736,6 +4832,9 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
if uErr == nil {
|
||||
if reqFlow < 0 {
|
||||
reqFlow = uFlow
|
||||
if user, err := h.repo.GetUserByID(userID); err == nil && user != nil {
|
||||
reqFlowMiB = user.FlowMiB
|
||||
}
|
||||
}
|
||||
if reqNum < 0 {
|
||||
reqNum = uNum
|
||||
@@ -4764,7 +4863,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
reqStatus = 1
|
||||
}
|
||||
|
||||
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil {
|
||||
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus, reqFlowMiB); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -4788,8 +4887,20 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
}
|
||||
|
||||
newFlow := currentFlow
|
||||
oldTunnel, err := h.repo.GetUserTunnelByID(existingID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if oldTunnel == nil {
|
||||
return fmt.Errorf("隧道权限不存在")
|
||||
}
|
||||
newFlowMiB := oldTunnel.FlowMiB
|
||||
if reqFlow >= 0 {
|
||||
newFlow = reqFlow
|
||||
newFlowMiB = reqFlowMiB
|
||||
if _, supplied := req["flowMiB"]; !supplied && reqFlow == currentFlow {
|
||||
newFlowMiB = oldTunnel.FlowMiB
|
||||
}
|
||||
}
|
||||
|
||||
newNum := int(currentNum)
|
||||
@@ -4819,7 +4930,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
|
||||
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus, newFlowMiB); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -4832,6 +4943,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
currentExpTime,
|
||||
currentFlowReset,
|
||||
currentStatus,
|
||||
oldTunnel.FlowMiB,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr)
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -21,6 +22,7 @@ type nftablesRuntimeManager interface {
|
||||
Test(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error)
|
||||
Clear(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
CollectCounters(ctx context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error)
|
||||
}
|
||||
|
||||
func isNftablesForwardMode(mode string) bool {
|
||||
@@ -77,9 +79,13 @@ func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAdd
|
||||
if len(entryNodeIDs) != 1 {
|
||||
return errors.New("nftables 节点仅支持单入口隧道")
|
||||
}
|
||||
if _, err := runtimenft.ParseSingleTarget(remoteAddr); err != nil {
|
||||
target, err := runtimenft.ParseSingleTarget(remoteAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if net.ParseIP(strings.Trim(strings.TrimSpace(target.Host), "[]")) == nil {
|
||||
return errors.New("nftables 节点仅支持 IP 目标地址")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -20,21 +21,29 @@ import (
|
||||
)
|
||||
|
||||
type fakeNftablesManager struct {
|
||||
testErr error
|
||||
reconcileErr error
|
||||
reconcileHit int
|
||||
clearErr error
|
||||
clearHit int
|
||||
lastConfig runtimenft.SSHConfig
|
||||
lastPlan runtimenft.NodePlan
|
||||
mu sync.Mutex
|
||||
testErr error
|
||||
reconcileErr error
|
||||
reconcileHit int
|
||||
clearErr error
|
||||
clearHit int
|
||||
collectErr error
|
||||
collectHit int
|
||||
counterSamples []runtimenft.CounterSample
|
||||
lastConfig runtimenft.SSHConfig
|
||||
lastPlan runtimenft.NodePlan
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.lastConfig = cfg
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.reconcileHit++
|
||||
f.lastConfig = cfg
|
||||
f.lastPlan = plan
|
||||
@@ -44,15 +53,40 @@ func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHCon
|
||||
return runtimenft.ApplyResult{
|
||||
NodeID: plan.NodeID,
|
||||
Script: "table inet flvx {}",
|
||||
Hashes: map[int64]string{plan.NodeID: "hash"},
|
||||
Hashes: runtimenft.PlanHashes(plan),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.clearHit++
|
||||
return f.clearErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.collectHit++
|
||||
f.lastConfig = cfg
|
||||
if f.collectErr != nil {
|
||||
return nil, f.collectErr
|
||||
}
|
||||
return f.counterSamples, nil
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) reconcileCount() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return f.reconcileHit
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) collectCount() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return f.collectHit
|
||||
}
|
||||
|
||||
type nftablesTestFixture struct {
|
||||
handler *Handler
|
||||
nodeID int64
|
||||
@@ -146,6 +180,49 @@ func TestNodeNftablesReconcileEndpointPersistsBindings(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartBackgroundJobsReconcilesNftablesRulesAtStartup(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-startup-tunnel", fixture.nodeID)
|
||||
seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
h.StartBackgroundJobs()
|
||||
t.Cleanup(h.StopBackgroundJobs)
|
||||
|
||||
waitForCondition(t, time.Second, func() bool {
|
||||
return manager.reconcileCount() > 0
|
||||
}, "nftables startup reconcile")
|
||||
}
|
||||
|
||||
func TestStartBackgroundJobsCollectsNftablesTrafficImmediatelyAndUsesFastInterval(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}}
|
||||
h.nftablesManager = manager
|
||||
|
||||
oldInterval := nftablesTrafficCollectInterval
|
||||
nftablesTrafficCollectInterval = 20 * time.Millisecond
|
||||
t.Cleanup(func() { nftablesTrafficCollectInterval = oldInterval })
|
||||
|
||||
h.StartBackgroundJobs()
|
||||
t.Cleanup(h.StopBackgroundJobs)
|
||||
|
||||
waitForCondition(t, time.Second, func() bool {
|
||||
return manager.collectCount() >= 2
|
||||
}, "immediate and repeated nftables traffic collection")
|
||||
}
|
||||
|
||||
func TestNftablesTrafficCollectIntervalDefaultsToThirtySeconds(t *testing.T) {
|
||||
if nftablesTrafficCollectInterval != 30*time.Second {
|
||||
t.Fatalf("expected default nftables traffic collection interval 30s, got %s", nftablesTrafficCollectInterval)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesClearEndpointClearsBindings(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
@@ -268,6 +345,83 @@ func TestNodeUpdatePreservesExistingNftablesSecretsWhenFieldsOmitted(t *testing.
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeUpdateSkipsAgentProtocolCommandForNftablesNode(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
|
||||
"id": fixture.nodeID,
|
||||
"name": "nft-node-updated",
|
||||
"serverIp": "198.51.100.10",
|
||||
"serverIpV4": "198.51.100.10",
|
||||
"port": "1000-65535",
|
||||
"forwardMode": "nftables",
|
||||
"http": 1,
|
||||
"tls": 1,
|
||||
"socks": 1,
|
||||
"sshConfig": map[string]interface{}{
|
||||
"host": "203.0.113.30",
|
||||
"port": 22,
|
||||
"username": "admin",
|
||||
"authType": "password",
|
||||
"sudoMode": "none",
|
||||
},
|
||||
})
|
||||
res := httptest.NewRecorder()
|
||||
fixture.handler.nodeUpdate(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
|
||||
cfg, err := fixture.handler.repo.GetNodeSSHConfig(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
if !cfg.Password.Valid || cfg.Password.String != "secret" {
|
||||
t.Fatalf("expected password secret to be preserved, got %+v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateNftablesForwardRequestRejectsHostnameTarget(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
tunnel, err := h.getTunnelRecord(tunnelID)
|
||||
if err != nil {
|
||||
t.Fatalf("load tunnel: %v", err)
|
||||
}
|
||||
|
||||
err = h.validateNftablesForwardRequest(tunnel, "example.com:443", []int64{fixture.nodeID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected hostname target to be rejected")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "IP") {
|
||||
t.Fatalf("expected IP literal validation error, got %q", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardDeleteReconcilesNftablesAfterDBDelete(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]int64{"id": forward.ID})
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardDelete(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if len(manager.lastPlan.Rules) != 0 {
|
||||
t.Fatalf("expected reconcile after DB delete to render no rules, got %+v", manager.lastPlan.Rules)
|
||||
}
|
||||
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
|
||||
t.Fatalf("expected forward to be deleted, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
@@ -307,6 +461,30 @@ func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardBatchDeleteReconcilesNftablesAfterDBDelete(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {forward.ID}})
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardBatchDelete(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if len(manager.lastPlan.Rules) != 0 {
|
||||
t.Fatalf("expected reconcile after DB delete to render no rules, got %+v", manager.lastPlan.Rules)
|
||||
}
|
||||
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
|
||||
t.Fatalf("expected forward to be deleted, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
@@ -343,6 +521,61 @@ func TestTunnelBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiagnoseForwardRuntimeReturnsNftablesRuleStatus(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-diagnose-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forward.ID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
RuleHash: "hash-a",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, now); err != nil {
|
||||
t.Fatalf("seed nft binding: %v", err)
|
||||
}
|
||||
|
||||
payload, err := h.diagnoseForwardRuntime(context.Background(), &forwardRecord{
|
||||
ID: forward.ID,
|
||||
Name: forward.Name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: "203.0.113.9:8080",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("diagnose forward: %v", err)
|
||||
}
|
||||
results, ok := payload["results"].([]map[string]interface{})
|
||||
if !ok || len(results) != 1 {
|
||||
t.Fatalf("expected one nftables diagnosis result, got %#v", payload["results"])
|
||||
}
|
||||
result := results[0]
|
||||
if result["forwardMode"] != "nftables" || result["nftRuleStatus"] != runtimenft.StatusApplied {
|
||||
t.Fatalf("expected nftables applied result, got %#v", result)
|
||||
}
|
||||
if result["success"] != true {
|
||||
t.Fatalf("expected nftables diagnosis success, got %#v", result)
|
||||
}
|
||||
if !strings.Contains(asString(result["message"]), "已下发") {
|
||||
t.Fatalf("expected applied message, got %#v", result["message"])
|
||||
}
|
||||
}
|
||||
|
||||
func waitForCondition(t *testing.T, timeout time.Duration, condition func() bool, description string) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
if condition() {
|
||||
return
|
||||
}
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("timed out waiting for %s", description)
|
||||
}
|
||||
|
||||
func setupNftablesHandler(t *testing.T) nftablesTestFixture {
|
||||
t.Helper()
|
||||
|
||||
@@ -412,7 +645,7 @@ func seedForwardForNftables(t *testing.T, h *Handler, tunnelID, nodeID int64, re
|
||||
now := time.Now().UnixMilli()
|
||||
forwardID, err := h.repo.CreateForwardTx(
|
||||
1, "admin", "nft-forward", tunnelID, remoteAddr, "fifo", now, 1,
|
||||
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0,
|
||||
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0, 0, 0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create forward: %v", err)
|
||||
|
||||
@@ -0,0 +1,400 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"math"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type nftTrafficDelta struct {
|
||||
ForwardID int64
|
||||
BytesIn int64
|
||||
BytesOut int64
|
||||
}
|
||||
|
||||
type nftCounterStateKey struct {
|
||||
forwardID int64
|
||||
protocol string
|
||||
direction string
|
||||
}
|
||||
|
||||
func (h *Handler) runNftablesTrafficCollectJob(now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
nodes, err := h.repo.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=list_nodes err=%v", err)
|
||||
return
|
||||
}
|
||||
for i := range nodes {
|
||||
node := &nodes[i]
|
||||
h.collectNftablesNodeTraffic(node.NodeID, &node.Config, now)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) collectNftablesNodeTraffic(nodeID int64, cfgModel *model.NodeSSHConfig, now time.Time) {
|
||||
if h == nil || h.repo == nil {
|
||||
return
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
log.Printf("nftables traffic collection failed op=collect node_id=%d err=%v", nodeID, "nftables manager not initialized")
|
||||
return
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfgModel)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=ssh_config node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
samples, err := h.nftablesManager.CollectCounters(context.Background(), sshCfg)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=collect node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
oldStates, err := h.repo.GetNftCounterStatesByNode(nodeID)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=list_states node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=list_bindings node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
hashes := make(map[int64]string, len(bindings))
|
||||
for _, binding := range bindings {
|
||||
if strings.ToLower(strings.TrimSpace(binding.Status)) != runtimenft.StatusApplied {
|
||||
continue
|
||||
}
|
||||
ruleHash := strings.TrimSpace(binding.RuleHash)
|
||||
if ruleHash == "" {
|
||||
continue
|
||||
}
|
||||
hashes[binding.ForwardID] = ruleHash
|
||||
}
|
||||
|
||||
nowMs := now.UnixMilli()
|
||||
boundSamples := filterNftCounterSamplesWithBinding(samples, hashes)
|
||||
deltas, newStates := buildNftCounterDeltas(nodeID, boundSamples, oldStates, hashes, nowMs)
|
||||
if len(newStates) == 0 {
|
||||
if len(deltas) != 0 {
|
||||
log.Printf("nftables traffic collection skipped suspicious deltas without states node_id=%d deltas=%d", nodeID, len(deltas))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
var metas map[int64]repo.FlowUploadForwardMeta
|
||||
forwardIDs := make([]int64, 0, len(deltas))
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID > 0 {
|
||||
forwardIDs = append(forwardIDs, delta.ForwardID)
|
||||
}
|
||||
}
|
||||
if len(deltas) != 0 {
|
||||
metas, err = h.repo.GetFlowUploadForwardMetas(forwardIDs)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=load_flow_metas node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
if missingForwardID, ok := firstNftDeltaMissingMeta(deltas, metas); ok {
|
||||
log.Printf("nftables traffic collection skipped state advance op=missing_flow_meta node_id=%d forward_id=%d", nodeID, missingForwardID)
|
||||
return
|
||||
}
|
||||
}
|
||||
if len(deltas) == 0 {
|
||||
if err := h.repo.UpsertNftCounterStates(newStates, nowMs); err != nil {
|
||||
log.Printf("nftables traffic collection failed op=upsert_states node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
batch := buildNftFlowUploadBatch(deltas, metas)
|
||||
if missingForwardID, ok := firstNftBatchMissingDelta(deltas, batch); ok {
|
||||
log.Printf("nftables traffic collection skipped state advance op=unaccounted_delta node_id=%d forward_id=%d", nodeID, missingForwardID)
|
||||
return
|
||||
}
|
||||
quotaViews, err := h.repo.ApplyNftTrafficAccounting(batch.flowDeltas, batch.quotaUsage, newStates, now)
|
||||
if err != nil {
|
||||
log.Printf("nftables traffic collection failed op=accounting node_id=%d err=%v", nodeID, err)
|
||||
return
|
||||
}
|
||||
h.recordTunnelMetricsFromForwardBatch(nodeID, batch.forwardTraffic, metas, nowMs)
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
func firstNftBatchMissingDelta(deltas []nftTrafficDelta, batch flowUploadBatch) (int64, bool) {
|
||||
flowSeen := make(map[int64]struct{}, len(batch.flowDeltas))
|
||||
for _, delta := range batch.flowDeltas {
|
||||
flowSeen[delta.ForwardID] = struct{}{}
|
||||
}
|
||||
|
||||
expectedRaw := make(map[int64]tunnelTrafficDelta, len(batch.forwardTraffic))
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID <= 0 || (delta.BytesIn == 0 && delta.BytesOut == 0) {
|
||||
continue
|
||||
}
|
||||
if delta.BytesIn < 0 || delta.BytesOut < 0 {
|
||||
return delta.ForwardID, true
|
||||
}
|
||||
raw := expectedRaw[delta.ForwardID]
|
||||
if raw.bytesIn > math.MaxInt64-delta.BytesIn || raw.bytesOut > math.MaxInt64-delta.BytesOut {
|
||||
return delta.ForwardID, true
|
||||
}
|
||||
raw.bytesIn += delta.BytesIn
|
||||
raw.bytesOut += delta.BytesOut
|
||||
expectedRaw[delta.ForwardID] = raw
|
||||
}
|
||||
|
||||
for forwardID, expected := range expectedRaw {
|
||||
actual, ok := batch.forwardTraffic[forwardID]
|
||||
if !ok || actual.bytesIn != expected.bytesIn || actual.bytesOut != expected.bytesOut {
|
||||
return forwardID, true
|
||||
}
|
||||
if expected.bytesIn != 0 || expected.bytesOut != 0 {
|
||||
if _, ok := flowSeen[forwardID]; !ok {
|
||||
return forwardID, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func firstNftDeltaMissingMeta(deltas []nftTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta) (int64, bool) {
|
||||
for _, delta := range deltas {
|
||||
if delta.ForwardID <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := metas[delta.ForwardID]; !ok {
|
||||
return delta.ForwardID, true
|
||||
}
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func filterNftCounterSamplesWithBinding(samples []runtimenft.CounterSample, hashes map[int64]string) []runtimenft.CounterSample {
|
||||
if len(samples) == 0 || len(hashes) == 0 {
|
||||
return nil
|
||||
}
|
||||
filtered := make([]runtimenft.CounterSample, 0, len(samples))
|
||||
for _, sample := range samples {
|
||||
if _, ok := hashes[sample.ForwardID]; !ok {
|
||||
continue
|
||||
}
|
||||
filtered = append(filtered, sample)
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func nftCounterKey(forwardID int64, protocol, direction string) nftCounterStateKey {
|
||||
return nftCounterStateKey{
|
||||
forwardID: forwardID,
|
||||
protocol: strings.ToLower(strings.TrimSpace(protocol)),
|
||||
direction: strings.ToLower(strings.TrimSpace(direction)),
|
||||
}
|
||||
}
|
||||
|
||||
func buildNftCounterDeltas(nodeID int64, samples []runtimenft.CounterSample, oldStates []model.NftCounterState, hashes map[int64]string, nowMs int64) ([]nftTrafficDelta, []repo.NftCounterStateInput) {
|
||||
oldByKey := make(map[nftCounterStateKey]model.NftCounterState, len(oldStates))
|
||||
for _, old := range oldStates {
|
||||
if old.NodeID != nodeID {
|
||||
continue
|
||||
}
|
||||
oldByKey[nftCounterKey(old.ForwardID, old.Protocol, old.Direction)] = old
|
||||
}
|
||||
|
||||
stateInputs := make([]repo.NftCounterStateInput, 0, len(samples))
|
||||
deltaByForward := make(map[int64]nftTrafficDelta)
|
||||
for _, sample := range samples {
|
||||
direction := strings.ToLower(strings.TrimSpace(sample.Direction))
|
||||
if direction != runtimenft.CounterDirectionToTarget && direction != runtimenft.CounterDirectionFromTarget {
|
||||
continue
|
||||
}
|
||||
|
||||
protocol := strings.ToLower(strings.TrimSpace(sample.Protocol))
|
||||
if protocol != "tcp" && protocol != "udp" {
|
||||
continue
|
||||
}
|
||||
if sample.Bytes > uint64(math.MaxInt64) || sample.Packets > uint64(math.MaxInt64) {
|
||||
continue
|
||||
}
|
||||
ruleHash := strings.TrimSpace(hashes[sample.ForwardID])
|
||||
stateInput := repo.NftCounterStateInput{
|
||||
NodeID: nodeID,
|
||||
ForwardID: sample.ForwardID,
|
||||
Protocol: protocol,
|
||||
Direction: direction,
|
||||
RuleHash: ruleHash,
|
||||
Bytes: sample.Bytes,
|
||||
Packets: sample.Packets,
|
||||
CollectedTime: nowMs,
|
||||
}
|
||||
|
||||
old, exists := oldByKey[nftCounterKey(sample.ForwardID, protocol, direction)]
|
||||
if !exists || old.RuleHash != ruleHash {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
if old.Bytes < 0 {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
oldBytes := uint64(old.Bytes)
|
||||
if sample.Bytes < oldBytes {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
rawDelta := sample.Bytes - oldBytes
|
||||
if rawDelta == 0 {
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
continue
|
||||
}
|
||||
|
||||
delta := deltaByForward[sample.ForwardID]
|
||||
delta.ForwardID = sample.ForwardID
|
||||
rawDeltaInt := int64(rawDelta)
|
||||
if direction == runtimenft.CounterDirectionToTarget {
|
||||
if delta.BytesIn > math.MaxInt64-rawDeltaInt {
|
||||
continue
|
||||
}
|
||||
delta.BytesIn += rawDeltaInt
|
||||
} else {
|
||||
if delta.BytesOut > math.MaxInt64-rawDeltaInt {
|
||||
continue
|
||||
}
|
||||
delta.BytesOut += rawDeltaInt
|
||||
}
|
||||
stateInputs = append(stateInputs, stateInput)
|
||||
deltaByForward[sample.ForwardID] = delta
|
||||
}
|
||||
|
||||
forwardIDs := make([]int64, 0, len(deltaByForward))
|
||||
for forwardID := range deltaByForward {
|
||||
forwardIDs = append(forwardIDs, forwardID)
|
||||
}
|
||||
sort.Slice(forwardIDs, func(i, j int) bool { return forwardIDs[i] < forwardIDs[j] })
|
||||
|
||||
deltas := make([]nftTrafficDelta, 0, len(forwardIDs))
|
||||
for _, forwardID := range forwardIDs {
|
||||
delta := deltaByForward[forwardID]
|
||||
if delta.BytesIn == 0 && delta.BytesOut == 0 {
|
||||
continue
|
||||
}
|
||||
deltas = append(deltas, delta)
|
||||
}
|
||||
return deltas, stateInputs
|
||||
}
|
||||
|
||||
func buildNftFlowUploadBatch(deltas []nftTrafficDelta, 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 _, delta := range deltas {
|
||||
meta, exists := metas[delta.ForwardID]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
|
||||
raw := batch.forwardTraffic[delta.ForwardID]
|
||||
if delta.BytesIn < 0 || delta.BytesOut < 0 || raw.bytesIn > math.MaxInt64-delta.BytesIn || raw.bytesOut > math.MaxInt64-delta.BytesOut {
|
||||
continue
|
||||
}
|
||||
|
||||
scaledIn, ok := scaleNftTrafficBytes(delta.BytesIn, meta.TrafficRatio, meta.TunnelFlow)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
scaledOut, ok := scaleNftTrafficBytes(delta.BytesOut, meta.TrafficRatio, meta.TunnelFlow)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if scaledIn > math.MaxInt64-scaledOut {
|
||||
continue
|
||||
}
|
||||
quotaDelta := scaledIn + scaledOut
|
||||
if batch.quotaUsage[meta.UserID] > math.MaxInt64-quotaDelta {
|
||||
continue
|
||||
}
|
||||
|
||||
flowIdx, flowExists := flowSeen[delta.ForwardID]
|
||||
if flowExists && (batch.flowDeltas[flowIdx].InFlow > math.MaxInt64-scaledIn || batch.flowDeltas[flowIdx].OutFlow > math.MaxInt64-scaledOut) {
|
||||
continue
|
||||
}
|
||||
|
||||
raw.bytesIn += delta.BytesIn
|
||||
raw.bytesOut += delta.BytesOut
|
||||
batch.forwardTraffic[delta.ForwardID] = raw
|
||||
|
||||
if flowExists {
|
||||
batch.flowDeltas[flowIdx].InFlow += scaledIn
|
||||
batch.flowDeltas[flowIdx].OutFlow += scaledOut
|
||||
} else {
|
||||
flowSeen[delta.ForwardID] = len(batch.flowDeltas)
|
||||
batch.flowDeltas = append(batch.flowDeltas, repo.FlowUploadCounterDelta{
|
||||
ForwardID: delta.ForwardID,
|
||||
UserID: meta.UserID,
|
||||
UserTunnelID: meta.UserTunnelID,
|
||||
InFlow: scaledIn,
|
||||
OutFlow: scaledOut,
|
||||
})
|
||||
}
|
||||
batch.quotaUsage[meta.UserID] += quotaDelta
|
||||
|
||||
target := flowPolicyTarget{UserID: meta.UserID, UserTunnelID: meta.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 scaleNftTrafficBytes(bytes int64, ratio float64, tunnelFlow int64) (int64, bool) {
|
||||
if bytes < 0 || ratio < 0 || tunnelFlow < 0 {
|
||||
return 0, false
|
||||
}
|
||||
var scaled int64
|
||||
if ratio == 1 {
|
||||
scaled = bytes
|
||||
} else {
|
||||
scaledFloat := float64(bytes) * ratio
|
||||
if math.IsNaN(scaledFloat) || math.IsInf(scaledFloat, 0) || scaledFloat < 0 || scaledFloat >= math.Pow(2, 63) {
|
||||
return 0, false
|
||||
}
|
||||
scaled = int64(scaledFloat)
|
||||
}
|
||||
if tunnelFlow != 0 && scaled > math.MaxInt64/tunnelFlow {
|
||||
return 0, false
|
||||
}
|
||||
return scaled * tunnelFlow, true
|
||||
}
|
||||
@@ -0,0 +1,680 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildNftCounterDeltasSavesFirstBaselineWithoutDelta(t *testing.T) {
|
||||
nowMs := int64(1700000000123)
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}, nil, map[int64]string{42: "hash-a"}, nowMs)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected no deltas for first baseline, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one state input, got %d", len(states))
|
||||
}
|
||||
state := states[0]
|
||||
if state.NodeID != 11 || state.ForwardID != 42 || state.Protocol != "tcp" || state.Direction != runtimenft.CounterDirectionToTarget {
|
||||
t.Fatalf("unexpected state identity: %#v", state)
|
||||
}
|
||||
if state.RuleHash != "hash-a" || state.Bytes != 1000 || state.Packets != 10 || state.CollectedTime != nowMs {
|
||||
t.Fatalf("unexpected state values: %#v", state)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasNormalGrowthProducesDirectionalBytes(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: 15},
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2600, Packets: 26},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, RuleHash: "hash-a", Bytes: 2000, Packets: 20},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 1 {
|
||||
t.Fatalf("expected one aggregated delta, got %#v", deltas)
|
||||
}
|
||||
if deltas[0].ForwardID != 42 || deltas[0].BytesIn != 500 || deltas[0].BytesOut != 600 {
|
||||
t.Fatalf("unexpected delta: %#v", deltas[0])
|
||||
}
|
||||
if len(states) != 2 {
|
||||
t.Fatalf("expected two state inputs, got %d", len(states))
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasResetRefreshesBaselineWithoutDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 25, Packets: 2},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected reset to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 || states[0].Bytes != 25 || states[0].RuleHash != "hash-a" {
|
||||
t.Fatalf("expected refreshed baseline state, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasRuleHashChangeRefreshesBaselineWithoutDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: 15},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-b"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected rule hash change to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 || states[0].Bytes != 1500 || states[0].RuleHash != "hash-b" {
|
||||
t.Fatalf("expected refreshed hash baseline state, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasEqualBytesRefreshesBaselineWithoutDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 11},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected equal bytes to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 1 || states[0].Bytes != 1000 || states[0].Packets != 11 || states[0].RuleHash != "hash-a" {
|
||||
t.Fatalf("expected refreshed baseline state, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasAggregatesProtocolsAndDirections(t *testing.T) {
|
||||
deltas, _ := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1100, Packets: 11},
|
||||
{ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 2200, Packets: 22},
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 3300, Packets: 33},
|
||||
{ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 4400, Packets: 44},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 2000},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, RuleHash: "hash-a", Bytes: 3000},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionFromTarget, RuleHash: "hash-a", Bytes: 4000},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 1 {
|
||||
t.Fatalf("expected one aggregated delta, got %#v", deltas)
|
||||
}
|
||||
if deltas[0].ForwardID != 42 || deltas[0].BytesIn != 300 || deltas[0].BytesOut != 700 {
|
||||
t.Fatalf("unexpected aggregated delta: %#v", deltas[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsInvalidProtocolBeforeStateAndDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "icmp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: 15},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "icmp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected invalid protocol to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected invalid protocol to produce no state inputs, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsOversizedPacketsBeforeStateAndDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1500, Packets: uint64(math.MaxInt64) + 1},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected oversized packets to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected oversized packets to produce no state inputs, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsOversizedBytesBeforeStateAndDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64) + 1, Packets: 10},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1000, Packets: 10},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected oversized bytes to produce no deltas, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected oversized bytes to produce no state inputs, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsOverflowingAggregateSampleWithoutState(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64), Packets: 10},
|
||||
{ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 10, Packets: 1},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1, Packets: 1},
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "udp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-a", Bytes: 1, Packets: 1},
|
||||
}, map[int64]string{42: "hash-a"}, 2000)
|
||||
|
||||
if len(deltas) != 1 {
|
||||
t.Fatalf("expected only non-overflowing aggregate delta, got %#v", deltas)
|
||||
}
|
||||
if deltas[0].ForwardID != 42 || deltas[0].BytesIn != math.MaxInt64-1 || deltas[0].BytesOut != 0 {
|
||||
t.Fatalf("unexpected aggregate delta: %#v", deltas[0])
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected only the accounted safe sample to advance baseline, got %#v", states)
|
||||
}
|
||||
if states[0].ForwardID != 42 || states[0].Protocol != "tcp" || states[0].Bytes != uint64(math.MaxInt64) {
|
||||
t.Fatalf("expected safe sample state input to be preserved, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftCounterDeltasSkipsUnknownDirectionAndOversizedDelta(t *testing.T) {
|
||||
deltas, states := buildNftCounterDeltas(11, []runtimenft.CounterSample{
|
||||
{ForwardID: 42, Protocol: "tcp", Direction: "sideways", Bytes: 1500, Packets: 15},
|
||||
{ForwardID: 43, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64) + 1, Packets: 1},
|
||||
}, []model.NftCounterState{
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, RuleHash: "hash-b", Bytes: 100},
|
||||
}, map[int64]string{42: "hash-a", 43: "hash-b"}, 2000)
|
||||
|
||||
if len(deltas) != 0 {
|
||||
t.Fatalf("expected no delta for skipped/oversized samples, got %#v", deltas)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state inputs for skipped/oversized samples, got %#v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchScalesFlowAndPreservesRawTunnelTraffic(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: 80, BytesOut: 110},
|
||||
{ForwardID: 21, BytesIn: 7, BytesOut: 11},
|
||||
{ForwardID: 20, BytesIn: 20, BytesOut: 10},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TunnelID: 1, TrafficRatio: 2, TunnelFlow: 3},
|
||||
21: {ForwardID: 21, UserID: 2, UserTunnelID: 10, TunnelID: 1, TrafficRatio: 1.5, TunnelFlow: 2},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 2 {
|
||||
t.Fatalf("expected two flow deltas, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if batch.flowDeltas[0].ForwardID != 20 || batch.flowDeltas[0].InFlow != 600 || batch.flowDeltas[0].OutFlow != 720 {
|
||||
t.Fatalf("unexpected first flow delta: %#v", batch.flowDeltas[0])
|
||||
}
|
||||
if batch.flowDeltas[1].ForwardID != 21 || batch.flowDeltas[1].InFlow != 20 || batch.flowDeltas[1].OutFlow != 32 {
|
||||
t.Fatalf("unexpected second flow delta: %#v", batch.flowDeltas[1])
|
||||
}
|
||||
if batch.quotaUsage[2] != 1372 {
|
||||
t.Fatalf("expected quota usage 1372, got %d", batch.quotaUsage[2])
|
||||
}
|
||||
if len(batch.policyTargets) != 1 || batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("expected deduped policy target, got %#v", batch.policyTargets)
|
||||
}
|
||||
if traffic := batch.forwardTraffic[20]; traffic.bytesIn != 100 || traffic.bytesOut != 120 {
|
||||
t.Fatalf("expected raw traffic for forward 20, got %#v", traffic)
|
||||
}
|
||||
if traffic := batch.forwardTraffic[21]; traffic.bytesIn != 7 || traffic.bytesOut != 11 {
|
||||
t.Fatalf("expected raw traffic for forward 21, got %#v", traffic)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsOverflowingScaledFlow(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 2, TunnelFlow: 2},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 0 {
|
||||
t.Fatalf("expected overflowing scaled flow to be skipped, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if len(batch.quotaUsage) != 0 {
|
||||
t.Fatalf("expected no quota usage for overflowing scaled flow, got %#v", batch.quotaUsage)
|
||||
}
|
||||
if len(batch.policyTargets) != 0 {
|
||||
t.Fatalf("expected no policy targets for overflowing scaled flow, got %#v", batch.policyTargets)
|
||||
}
|
||||
if len(batch.forwardTraffic) != 0 {
|
||||
t.Fatalf("expected no raw traffic for overflowing scaled flow, got %#v", batch.forwardTraffic)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsRawForwardTrafficOverflow(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
{ForwardID: 20, BytesIn: 1, BytesOut: 0},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 0.5, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
traffic := batch.forwardTraffic[20]
|
||||
if traffic.bytesIn != math.MaxInt64 || traffic.bytesOut != 0 {
|
||||
t.Fatalf("expected overflowing raw delta to be skipped without negative traffic, got %#v", traffic)
|
||||
}
|
||||
if len(batch.flowDeltas) != 1 || batch.flowDeltas[0].ForwardID != 20 {
|
||||
t.Fatalf("expected only the safe flow delta, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if len(batch.policyTargets) != 1 || batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("expected policy target only from safe delta, got %#v", batch.policyTargets)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsQuotaOverflow(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
{ForwardID: 21, BytesIn: 1, BytesOut: 0},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 1, TunnelFlow: 1},
|
||||
21: {ForwardID: 21, UserID: 2, UserTunnelID: 10, TrafficRatio: 1, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 1 || batch.flowDeltas[0].ForwardID != 20 || batch.flowDeltas[0].InFlow != math.MaxInt64 {
|
||||
t.Fatalf("expected only non-overflowing quota delta, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if batch.quotaUsage[2] != math.MaxInt64 {
|
||||
t.Fatalf("expected quota usage to remain at max int64, got %#v", batch.quotaUsage)
|
||||
}
|
||||
if len(batch.policyTargets) != 1 || batch.policyTargets[0].UserID != 2 || batch.policyTargets[0].UserTunnelID != 10 {
|
||||
t.Fatalf("expected one policy target from non-overflowing delta, got %#v", batch.policyTargets)
|
||||
}
|
||||
if _, ok := batch.forwardTraffic[21]; ok {
|
||||
t.Fatalf("expected quota-overflowing delta to be skipped from raw traffic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildNftFlowUploadBatchSkipsMissingMeta(t *testing.T) {
|
||||
batch := buildNftFlowUploadBatch([]nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: 80, BytesOut: 110},
|
||||
{ForwardID: 99, BytesIn: 1, BytesOut: 2},
|
||||
}, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 1, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
if len(batch.flowDeltas) != 1 || batch.flowDeltas[0].ForwardID != 20 {
|
||||
t.Fatalf("expected only forward 20 delta, got %#v", batch.flowDeltas)
|
||||
}
|
||||
if _, ok := batch.forwardTraffic[99]; ok {
|
||||
t.Fatalf("expected missing meta forward to be skipped from raw traffic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftBatchCoversDeltasRequiresRawAndFlowEntries(t *testing.T) {
|
||||
deltas := []nftTrafficDelta{{ForwardID: 20, BytesIn: 1, BytesOut: 0}}
|
||||
batch := flowUploadBatch{
|
||||
forwardTraffic: map[int64]tunnelTrafficDelta{20: {bytesIn: 1}},
|
||||
flowDeltas: []repo.FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 1}},
|
||||
}
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); ok || missing != 0 {
|
||||
t.Fatalf("expected batch to cover delta, missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
|
||||
delete(batch.forwardTraffic, 20)
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
|
||||
t.Fatalf("expected missing raw traffic for forward 20, got missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
|
||||
batch.forwardTraffic[20] = tunnelTrafficDelta{bytesIn: 1}
|
||||
batch.flowDeltas = nil
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
|
||||
t.Fatalf("expected missing flow delta for forward 20, got missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftBatchCoversDeltasRequiresAggregateRawTotals(t *testing.T) {
|
||||
deltas := []nftTrafficDelta{
|
||||
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
|
||||
{ForwardID: 20, BytesIn: 1, BytesOut: 0},
|
||||
}
|
||||
batch := buildNftFlowUploadBatch(deltas, map[int64]repo.FlowUploadForwardMeta{
|
||||
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 0.5, TunnelFlow: 1},
|
||||
})
|
||||
|
||||
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
|
||||
t.Fatalf("expected aggregate raw overflow/mismatch for forward 20, got missing=%d ok=%v", missing, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficFirstBaselineSavesStateWithoutFlow(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2000, Packets: 20},
|
||||
}}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
if manager.collectHit != 1 {
|
||||
t.Fatalf("expected one collection, got %d", manager.collectHit)
|
||||
}
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 2 {
|
||||
t.Fatalf("expected two baseline states, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no forward flow on baseline, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT out_flow FROM user WHERE id = 1`); got != 0 {
|
||||
t.Fatalf("expected no user flow on baseline, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficGrowthAppliesFlowAndUpdatesState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2000, Packets: 20},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2600, Packets: 26},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 400 {
|
||||
t.Fatalf("expected forward in_flow=400, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT out_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 600 {
|
||||
t.Fatalf("expected forward out_flow=600, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM user WHERE id = 1`); got != 400 {
|
||||
t.Fatalf("expected user in_flow=400, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT out_flow FROM user_tunnel WHERE id = ?`, fixture.userTunnelID); got != 600 {
|
||||
t.Fatalf("expected user_tunnel out_flow=600, got %d", got)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 1), 0)`); got != 1000 {
|
||||
t.Fatalf("expected daily quota usage=1000, got %d", got)
|
||||
}
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 2 {
|
||||
t.Fatalf("expected two states after growth, got %+v", states)
|
||||
}
|
||||
for _, state := range states {
|
||||
if state.Direction == runtimenft.CounterDirectionToTarget && state.Bytes != 1400 {
|
||||
t.Fatalf("expected to-target state bytes 1400, got %+v", state)
|
||||
}
|
||||
if state.Direction == runtimenft.CounterDirectionFromTarget && state.Bytes != 2600 {
|
||||
t.Fatalf("expected from-target state bytes 2600, got %+v", state)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficSkippedBatchDeltaDoesNotAdvanceState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
if err := h.repo.DB().Exec(`UPDATE tunnel SET traffic_ratio = 2 WHERE id = (SELECT tunnel_id FROM forward WHERE id = ?)`, fixture.forwardID).Error; err != nil {
|
||||
t.Fatalf("update tunnel ratio: %v", err)
|
||||
}
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 0, Packets: 0},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64), Packets: 1},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one state, got %+v", states)
|
||||
}
|
||||
if states[0].Bytes != 0 || states[0].Packets != 0 {
|
||||
t.Fatalf("expected state to remain at old baseline after skipped batch delta, got %+v", states[0])
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no forward flow for skipped batch delta, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficMetadataErrorDoesNotAdvanceState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
if err := h.repo.DB().Exec(`DROP TABLE tunnel`).Error; err != nil {
|
||||
t.Fatalf("drop tunnel table: %v", err)
|
||||
}
|
||||
manager.counterSamples = []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
|
||||
}
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one baseline state, got %+v", states)
|
||||
}
|
||||
if states[0].Bytes != 1000 || states[0].Packets != 10 {
|
||||
t.Fatalf("expected state to remain at first baseline after metadata failure, got %+v", states[0])
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow after metadata failure, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficMissingMetaDoesNotAdvanceState(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
forwardID := int64(4242)
|
||||
nowMs := time.Now().UnixMilli()
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forwardID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
RuleHash: "hash-a",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, nowMs); err != nil {
|
||||
t.Fatalf("seed stale applied binding: %v", err)
|
||||
}
|
||||
if err := h.repo.UpsertNftCounterStates([]repo.NftCounterStateInput{{
|
||||
NodeID: fixture.nodeID,
|
||||
ForwardID: forwardID,
|
||||
Protocol: "tcp",
|
||||
Direction: runtimenft.CounterDirectionToTarget,
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1000,
|
||||
Packets: 10,
|
||||
CollectedTime: nowMs,
|
||||
}}, nowMs); err != nil {
|
||||
t.Fatalf("seed counter state: %v", err)
|
||||
}
|
||||
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
|
||||
}}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 1 {
|
||||
t.Fatalf("expected one state, got %+v", states)
|
||||
}
|
||||
if states[0].Bytes != 1000 || states[0].Packets != 10 {
|
||||
t.Fatalf("expected state to remain at old baseline when meta is missing, got %+v", states[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficSkipsSamplesWithoutBinding(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
if err := h.repo.DeleteNftRuleBindingsByForward(fixture.forwardID); err != nil {
|
||||
t.Fatalf("delete nft binding: %v", err)
|
||||
}
|
||||
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state for unbound sample, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow for unbound sample, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficSkipsNonAppliedBinding(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
if err := h.repo.MarkNftRuleBindingError(fixture.forwardID, fixture.nodeID, "apply failed", time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("mark binding error: %v", err)
|
||||
}
|
||||
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
|
||||
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
|
||||
}}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state for non-applied binding, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow for non-applied binding, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectNftablesNodeTrafficCollectionErrorDoesNotWriteState(t *testing.T) {
|
||||
fixture := setupNftablesCollectionFixture(t)
|
||||
h := fixture.handler
|
||||
h.nftablesManager = &fakeNftablesManager{collectErr: errors.New("ssh failed")}
|
||||
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
|
||||
|
||||
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
|
||||
|
||||
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load states: %v", err)
|
||||
}
|
||||
if len(states) != 0 {
|
||||
t.Fatalf("expected no state on collection error, got %+v", states)
|
||||
}
|
||||
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
|
||||
t.Fatalf("expected no flow on collection error, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
type nftablesCollectionFixture struct {
|
||||
handler *Handler
|
||||
nodeID int64
|
||||
forwardID int64
|
||||
userTunnelID int64
|
||||
}
|
||||
|
||||
func setupNftablesCollectionFixture(t *testing.T) nftablesCollectionFixture {
|
||||
t.Helper()
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-traffic-tunnel", fixture.nodeID)
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(1, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("seed user_tunnel: %v", err)
|
||||
}
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forward.ID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
RuleHash: "hash-a",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, now); err != nil {
|
||||
t.Fatalf("seed nft binding: %v", err)
|
||||
}
|
||||
userTunnelID := mustHandlerCount(t, h, `SELECT id FROM user_tunnel WHERE user_id = 1 AND tunnel_id = ?`, tunnelID)
|
||||
return nftablesCollectionFixture{
|
||||
handler: h,
|
||||
nodeID: fixture.nodeID,
|
||||
forwardID: forward.ID,
|
||||
userTunnelID: userTunnelID,
|
||||
}
|
||||
}
|
||||
|
||||
func mustCollectionSSHConfig(t *testing.T, h *Handler, nodeID int64) *model.NodeSSHConfig {
|
||||
t.Helper()
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
|
||||
func mustHandlerCount(t *testing.T, h *Handler, query string, args ...interface{}) int64 {
|
||||
t.Helper()
|
||||
var value int64
|
||||
if err := h.repo.DB().Raw(query, args...).Row().Scan(&value); err != nil {
|
||||
t.Fatalf("query %q failed: %v", query, err)
|
||||
}
|
||||
return value
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package handler
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
@@ -75,6 +76,12 @@ func IsValidNodeAddress(addr string) error {
|
||||
if strings.ContainsAny(addr, "/?") {
|
||||
return fmt.Errorf("address must not contain path or query parameters")
|
||||
}
|
||||
// A bare IPv6 literal contains multiple colons, so net.SplitHostPort treats
|
||||
// it as a malformed host:port pair. Accept IP literals before attempting
|
||||
// host:port parsing; netip also handles scoped IPv6 addresses.
|
||||
if _, err := netip.ParseAddr(addr); err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
_, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package handler
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestIssue515IsValidNodeAddressAcceptsBareIPv6(t *testing.T) {
|
||||
for _, addr := range []string{
|
||||
"2001:db8::1",
|
||||
"::1",
|
||||
"fe80::1%eth0",
|
||||
} {
|
||||
t.Run(addr, func(t *testing.T) {
|
||||
if err := IsValidNodeAddress(addr); err != nil {
|
||||
t.Fatalf("expected bare IPv6 address %q to be accepted: %v", addr, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsValidNodeAddressKeepsExistingAddressForms(t *testing.T) {
|
||||
for _, addr := range []string{
|
||||
"203.0.113.10",
|
||||
"node.example.com",
|
||||
"node.example.com:6365",
|
||||
"[2001:db8::1]:6365",
|
||||
} {
|
||||
t.Run(addr, func(t *testing.T) {
|
||||
if err := IsValidNodeAddress(addr); err != nil {
|
||||
t.Fatalf("expected node address %q to be accepted: %v", addr, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsValidNodeAddressRejectsURLComponents(t *testing.T) {
|
||||
for _, addr := range []string{
|
||||
"https://node.example.com",
|
||||
"node.example.com/path",
|
||||
"node.example.com?transport=tcp",
|
||||
} {
|
||||
t.Run(addr, func(t *testing.T) {
|
||||
if err := IsValidNodeAddress(addr); err == nil {
|
||||
t.Fatalf("expected node address %q to be rejected", addr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// flowMiB is optional so older clients can keep sending the GB-based flow field.
|
||||
// A positive value takes precedence and preserves sub-GB limits exactly.
|
||||
func parseTrafficLimit(req map[string]interface{}, defaultGB int64) (flowGB, flowMiB int64, err error) {
|
||||
flowGB = asInt64(req["flow"], defaultGB)
|
||||
if flowGB < 0 {
|
||||
return 0, 0, fmt.Errorf("流量限制不能小于0")
|
||||
}
|
||||
raw, present := req["flowMiB"]
|
||||
if !present {
|
||||
return flowGB, 0, nil
|
||||
}
|
||||
flowMiB, err = strconv.ParseInt(asString(raw), 10, 64)
|
||||
if err != nil || flowMiB < 0 || flowMiB > math.MaxInt64/bytesPerMiB {
|
||||
return 0, 0, fmt.Errorf("流量限制超出范围")
|
||||
}
|
||||
if flowMiB > 0 {
|
||||
flowGB = (flowMiB-1)/1024 + 1
|
||||
}
|
||||
return flowGB, flowMiB, nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestTrafficLimitMiBOverridesLegacyGB(t *testing.T) {
|
||||
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{
|
||||
"flow": float64(1), "flowMiB": float64(500),
|
||||
}, 100)
|
||||
if err != nil || flowGB != 1 || flowMiB != 500 {
|
||||
t.Fatalf("parseTrafficLimit = (%d, %d, %v), want (1, 500, nil)", flowGB, flowMiB, err)
|
||||
}
|
||||
limit := flowLimitBytes(flowGB, flowMiB)
|
||||
if limit != 500*bytesPerMiB {
|
||||
t.Fatalf("limit = %d, want %d", limit, 500*bytesPerMiB)
|
||||
}
|
||||
policy := &userTunnelPolicy{Flow: flowGB, FlowMiB: flowMiB, InFlow: limit - 1, Status: 1}
|
||||
if shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
|
||||
t.Fatal("policy paused before reaching 500 MiB")
|
||||
}
|
||||
policy.InFlow = limit
|
||||
if !shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
|
||||
t.Fatal("policy did not pause at 500 MiB")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrafficLimitLegacyAndInvalidValues(t *testing.T) {
|
||||
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{"flow": float64(2)}, 100)
|
||||
if err != nil || flowGB != 2 || flowMiB != 0 || flowLimitBytes(flowGB, flowMiB) != 2*bytesPerGB {
|
||||
t.Fatalf("legacy GB limit changed: (%d, %d, %v)", flowGB, flowMiB, err)
|
||||
}
|
||||
for _, value := range []interface{}{"1.5", -1, "999999999999999999999"} {
|
||||
if _, _, err := parseTrafficLimit(map[string]interface{}{"flowMiB": value}, 100); err == nil {
|
||||
t.Fatalf("accepted invalid flowMiB %v", value)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -152,7 +152,11 @@ func evaluateBestExitOwner(owner chainNodeRecord, exits []chainNodeRecord, nodes
|
||||
ownerNode := nodes[owner.NodeID]
|
||||
for _, exit := range exits {
|
||||
exitNode := nodes[exit.NodeID]
|
||||
if exitNode == nil {
|
||||
if !isTunnelProbeNodeOnline(ownerNode) {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "owner node offline"))
|
||||
continue
|
||||
}
|
||||
if !isTunnelProbeNodeOnline(exitNode) {
|
||||
scores = append(scores, failedBestExitCandidate(owner.NodeID, exit, "exit node unavailable"))
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -358,9 +358,9 @@ func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
|
||||
{NodeID: 31, NodeName: "exit-b", Port: 30031},
|
||||
}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
31: {ID: 31, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
|
||||
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
31: {ID: 31, Status: 1, ServerIP: "10.0.0.31", ServerIPv4: "10.0.0.31", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
switch {
|
||||
@@ -387,12 +387,30 @@ func TestEvaluateBestExitOwnerScoresAllCandidates(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerSkipsOfflineCandidate(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Status: 0, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
}
|
||||
ping := func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
t.Fatalf("offline best-exit candidate should not be probed: node=%d target=%s:%d", nodeID, ip, port)
|
||||
return 0, 100, nil
|
||||
}
|
||||
|
||||
scores := evaluateBestExitOwner(owner, exits, nodes, "", diagnosisExecOptions{}, defaultTunnelProbeTarget(), ping)
|
||||
if len(scores) != 1 || scores[0].Success {
|
||||
t.Fatalf("expected one failed offline candidate, got %+v", scores)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEvaluateBestExitOwnerUsesConfiguredPublicProbeTarget(t *testing.T) {
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry-a"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30001}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Name: "entry-a", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Name: "exit-a", ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
10: {ID: 10, Name: "entry-a", Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10"},
|
||||
30: {ID: 30, Name: "exit-a", Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30"},
|
||||
}
|
||||
target := tunnelProbeTarget{Host: "speed.example.com", Port: 8443}
|
||||
var calls []string
|
||||
@@ -419,8 +437,8 @@ func TestEvaluateBestExitOwnerMarksCandidateFailedWhenOwnerToExitFails(t *testin
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-a", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
10: {ID: 10, Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Status: 1, ServerIP: "10.0.0.30", ServerIPv4: "10.0.0.30", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
return 0, 100, errBestExitProbeForTest
|
||||
@@ -436,8 +454,8 @@ func TestEvaluateBestExitOwnerMarksCandidateFailedWhenTargetResolutionFails(t *t
|
||||
owner := chainNodeRecord{NodeID: 10, NodeName: "entry"}
|
||||
exits := []chainNodeRecord{{NodeID: 30, NodeName: "exit-v6", Port: 30030}}
|
||||
nodes := map[int64]*nodeRecord{
|
||||
10: {ID: 10, Name: "entry", ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Name: "exit-v6", ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
|
||||
10: {ID: 10, Name: "entry", Status: 1, ServerIP: "10.0.0.10", ServerIPv4: "10.0.0.10", TCPListenAddr: "[::]"},
|
||||
30: {ID: 30, Name: "exit-v6", Status: 1, ServerIP: "2001:db8::30", ServerIPv6: "2001:db8::30", TCPListenAddr: "[::]"},
|
||||
}
|
||||
pinger := func(nodeID int64, ip string, port int, _ diagnosisExecOptions) (float64, float64, error) {
|
||||
t.Fatalf("ping should not be called when target resolution fails: node=%d ip=%s port=%d", nodeID, ip, port)
|
||||
|
||||
@@ -3,6 +3,8 @@ package handler
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -13,7 +15,6 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
tunnelQualityProbeInterval = 1 * time.Second
|
||||
tunnelQualityProbeTimeout = 8 * time.Second
|
||||
tunnelQualityPingTimeoutMs = 5000
|
||||
tunnelQualityPruneInterval = 10 * time.Minute
|
||||
@@ -31,6 +32,26 @@ type TunnelQualityHop struct {
|
||||
TargetPort int `json:"targetPort,omitempty"`
|
||||
}
|
||||
|
||||
type TunnelQualityCandidateHop struct {
|
||||
TunnelQualityHop
|
||||
FromRole string `json:"fromRole"`
|
||||
ToRole string `json:"toRole"`
|
||||
HopIndex int `json:"hopIndex"`
|
||||
Selected bool `json:"selected"`
|
||||
ErrorMessage string `json:"errorMessage,omitempty"`
|
||||
}
|
||||
|
||||
type tunnelQualityChainDetails struct {
|
||||
PrimaryPath []TunnelQualityHop `json:"primaryPath,omitempty"`
|
||||
CandidateHops []TunnelQualityCandidateHop `json:"candidateHops,omitempty"`
|
||||
}
|
||||
|
||||
type tunnelQualityCandidateGroup struct {
|
||||
role string
|
||||
roleIndex int
|
||||
nodes []chainNodeRecord
|
||||
}
|
||||
|
||||
// tunnelQualitySnapshot is the in-memory latest probe result for a tunnel.
|
||||
type tunnelQualitySnapshot struct {
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
@@ -56,7 +77,7 @@ type tunnelQualityProber struct {
|
||||
cache sync.Map // tunnelID (int64) → *tunnelQualitySnapshot
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
interval time.Duration
|
||||
wake chan struct{}
|
||||
lastPrune int64
|
||||
probing int32 // atomic flag: 1 = probeAll running, 0 = idle
|
||||
probeNode bestExitProbeFunc
|
||||
@@ -65,8 +86,8 @@ type tunnelQualityProber struct {
|
||||
// newTunnelQualityProber creates a new prober (not yet running).
|
||||
func newTunnelQualityProber(h *Handler) *tunnelQualityProber {
|
||||
return &tunnelQualityProber{
|
||||
handler: h,
|
||||
interval: tunnelQualityProbeInterval,
|
||||
handler: h,
|
||||
wake: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -86,6 +107,16 @@ func (p *tunnelQualityProber) Stop() {
|
||||
p.cancel()
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) NotifyConfigChanged() {
|
||||
if p == nil || p.wake == nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case p.wake <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// GetAll returns all cached quality snapshots (latest per tunnel).
|
||||
func (p *tunnelQualityProber) GetAll() []tunnelQualitySnapshot {
|
||||
var items []tunnelQualitySnapshot
|
||||
@@ -109,20 +140,44 @@ func (p *tunnelQualityProber) loop() {
|
||||
// Run once immediately
|
||||
p.probeAll()
|
||||
|
||||
ticker := time.NewTicker(p.interval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
timer := time.NewTimer(p.probeInterval())
|
||||
select {
|
||||
case <-p.ctx.Done():
|
||||
stopAndDrainTunnelQualityTimer(timer)
|
||||
return
|
||||
case <-ticker.C:
|
||||
case <-p.wake:
|
||||
stopAndDrainTunnelQualityTimer(timer)
|
||||
continue
|
||||
case <-timer.C:
|
||||
p.probeAll()
|
||||
p.maybePrune()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func stopAndDrainTunnelQualityTimer(timer *time.Timer) {
|
||||
if timer == nil || timer.Stop() {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-timer.C:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeInterval() time.Duration {
|
||||
if p == nil || p.handler == nil || p.handler.repo == nil {
|
||||
return time.Duration(monitoring.DefaultTunnelQualityProbeIntervalSec) * time.Second
|
||||
}
|
||||
cfg, err := p.handler.repo.GetConfigsByNames([]string{monitoring.ConfigTunnelQualityProbeIntervalSec})
|
||||
if err != nil {
|
||||
return time.Duration(monitoring.DefaultTunnelQualityProbeIntervalSec) * time.Second
|
||||
}
|
||||
seconds := monitoring.TunnelQualityProbeIntervalSecondsFromConfigMap(cfg)
|
||||
return time.Duration(seconds) * time.Second
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) isEnabled() bool {
|
||||
if p == nil || p.handler == nil {
|
||||
return true
|
||||
@@ -246,15 +301,28 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
options := diagnosisExecOptions{
|
||||
commandTimeout: tunnelQualityProbeTimeout,
|
||||
pingTimeoutMS: tunnelQualityPingTimeoutMs,
|
||||
pingCount: 1,
|
||||
timeoutMessage: "探测超时",
|
||||
}
|
||||
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget)
|
||||
roundPinger := newBestExitRoundPinger(p.pingNode)
|
||||
p.probeBestExitOwners(tunnelID, inNodes, midNodesGrouped, outNodes, ipPreference, options, probeTarget, roundPinger)
|
||||
|
||||
entry, _, entryOnline := p.firstOnlineChainNode(inNodes)
|
||||
exit, _, exitOnline := p.firstOnlineChainNode(outNodes)
|
||||
selectedNodeIDs := make(map[string]int64, 2+len(midNodesGrouped))
|
||||
if entryOnline {
|
||||
selectedNodeIDs[tunnelQualityGroupKey("entry", 0)] = entry.NodeID
|
||||
}
|
||||
if exitOnline {
|
||||
selectedNodeIDs[tunnelQualityGroupKey("exit", 0)] = exit.NodeID
|
||||
}
|
||||
var primaryHops []TunnelQualityHop
|
||||
|
||||
switch tunnel.Type {
|
||||
case 1:
|
||||
// Port forwarding: entry → public probe target only.
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if entryOnline {
|
||||
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -262,24 +330,42 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
} else {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
} else {
|
||||
snap.ErrorMessage = "入口节点均不在线"
|
||||
}
|
||||
case 2:
|
||||
// Tunnel forwarding: entry → exit + exit → Bing
|
||||
probeOK := true
|
||||
|
||||
if len(inNodes) > 0 && len(outNodes) > 0 {
|
||||
var hops []TunnelQualityHop
|
||||
if !entryOnline {
|
||||
probeOK = false
|
||||
snap.ErrorMessage = "入口节点均不在线"
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
} else if !exitOnline {
|
||||
probeOK = false
|
||||
snap.ErrorMessage = "出口节点均不在线"
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
} else {
|
||||
var totalLat float64
|
||||
remainingSuccessProb := 1.0
|
||||
|
||||
nodesInPath := make([]chainNodeRecord, 0, 2+len(midNodesGrouped))
|
||||
nodesInPath = append(nodesInPath, inNodes[0])
|
||||
for _, midGroup := range midNodesGrouped {
|
||||
if len(midGroup) > 0 {
|
||||
nodesInPath = append(nodesInPath, midGroup[0])
|
||||
nodesInPath = append(nodesInPath, entry)
|
||||
for midIndex, midGroup := range midNodesGrouped {
|
||||
mid, _, online := p.firstOnlineChainNode(midGroup)
|
||||
if !online {
|
||||
probeOK = false
|
||||
snap.ErrorMessage = "中间节点组均不在线"
|
||||
break
|
||||
}
|
||||
nodesInPath = append(nodesInPath, mid)
|
||||
selectedNodeIDs[tunnelQualityGroupKey("middle", midIndex)] = mid.NodeID
|
||||
}
|
||||
if probeOK {
|
||||
nodesInPath = append(nodesInPath, exit)
|
||||
}
|
||||
nodesInPath = append(nodesInPath, outNodes[0])
|
||||
|
||||
for i := 0; i < len(nodesInPath)-1; i++ {
|
||||
source := nodesInPath[i]
|
||||
@@ -293,12 +379,12 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
}
|
||||
|
||||
targetNode, nodeErr := h.getNodeRecord(target.NodeID)
|
||||
if nodeErr != nil || targetNode == nil {
|
||||
if nodeErr != nil || !isTunnelProbeNodeOnline(targetNode) {
|
||||
snap.ErrorMessage = "节点 " + target.NodeName + " 不可用"
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
@@ -309,25 +395,25 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
break
|
||||
}
|
||||
|
||||
hop.TargetIP = targetIP
|
||||
hop.TargetPort = targetPort
|
||||
|
||||
lat, loss, err := p.pingNode(source.NodeID, targetIP, targetPort, options)
|
||||
lat, loss, err := roundPinger(source.NodeID, targetIP, targetPort, options)
|
||||
if err == nil {
|
||||
hop.Latency = lat
|
||||
hop.Loss = loss
|
||||
totalLat += lat
|
||||
remainingSuccessProb *= (1.0 - loss/100.0)
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
} else {
|
||||
probeOK = false
|
||||
hop.Latency = -1
|
||||
hop.Loss = 100
|
||||
hops = append(hops, hop)
|
||||
primaryHops = append(primaryHops, hop)
|
||||
if snap.ErrorMessage == "" {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
@@ -342,17 +428,11 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
snap.EntryToExitLatency = -1
|
||||
snap.EntryToExitLoss = 100
|
||||
}
|
||||
|
||||
if len(hops) > 0 {
|
||||
if b, err := json.Marshal(hops); err == nil {
|
||||
snap.ChainDetails = string(b)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Exit → Bing
|
||||
if len(outNodes) > 0 {
|
||||
lat, loss, err := p.pingNode(outNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if exitOnline {
|
||||
lat, loss, err := roundPinger(exit.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -367,8 +447,8 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
snap.Success = probeOK
|
||||
default:
|
||||
// Unknown type: entry → public probe target.
|
||||
if len(inNodes) > 0 {
|
||||
lat, loss, err := p.pingNode(inNodes[0].NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if entryOnline {
|
||||
lat, loss, err := roundPinger(entry.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if err == nil {
|
||||
snap.ExitToBingLatency = lat
|
||||
snap.ExitToBingLoss = loss
|
||||
@@ -376,13 +456,215 @@ func (p *tunnelQualityProber) probeTunnel(tunnelID int64) {
|
||||
} else {
|
||||
snap.ErrorMessage = err.Error()
|
||||
}
|
||||
} else {
|
||||
snap.ErrorMessage = "入口节点均不在线"
|
||||
}
|
||||
}
|
||||
|
||||
candidateHops := p.probeTunnelCandidateHops(
|
||||
tunnel.Type,
|
||||
inNodes,
|
||||
midNodesGrouped,
|
||||
outNodes,
|
||||
selectedNodeIDs,
|
||||
ipPreference,
|
||||
options,
|
||||
probeTarget,
|
||||
roundPinger,
|
||||
)
|
||||
if len(primaryHops) > 0 || len(candidateHops) > 0 {
|
||||
details := tunnelQualityChainDetails{
|
||||
PrimaryPath: primaryHops,
|
||||
CandidateHops: candidateHops,
|
||||
}
|
||||
if b, err := json.Marshal(details); err == nil {
|
||||
snap.ChainDetails = string(b)
|
||||
}
|
||||
}
|
||||
|
||||
p.storeResult(snap)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget) {
|
||||
func tunnelQualityGroupKey(role string, index int) string {
|
||||
return fmt.Sprintf("%s:%d", role, index)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeTunnelCandidateHops(
|
||||
tunnelType int,
|
||||
inNodes []chainNodeRecord,
|
||||
chainHops [][]chainNodeRecord,
|
||||
outNodes []chainNodeRecord,
|
||||
selectedNodeIDs map[string]int64,
|
||||
ipPreference string,
|
||||
options diagnosisExecOptions,
|
||||
probeTarget tunnelProbeTarget,
|
||||
ping bestExitProbeFunc,
|
||||
) []TunnelQualityCandidateHop {
|
||||
if p == nil || p.handler == nil || ping == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if tunnelType != 2 {
|
||||
return p.probePublicTargetCandidates("entry", 0, inNodes, selectedNodeIDs, options, probeTarget, ping)
|
||||
}
|
||||
|
||||
groups := make([]tunnelQualityCandidateGroup, 0, 2+len(chainHops))
|
||||
groups = append(groups, tunnelQualityCandidateGroup{role: "entry", roleIndex: 0, nodes: inNodes})
|
||||
for i, hop := range chainHops {
|
||||
groups = append(groups, tunnelQualityCandidateGroup{role: "middle", roleIndex: i, nodes: hop})
|
||||
}
|
||||
groups = append(groups, tunnelQualityCandidateGroup{role: "exit", roleIndex: 0, nodes: outNodes})
|
||||
|
||||
var items []TunnelQualityCandidateHop
|
||||
for i := 0; i < len(groups)-1; i++ {
|
||||
items = append(items, p.probeCandidateGroupLinks(
|
||||
groups[i],
|
||||
groups[i+1],
|
||||
i,
|
||||
selectedNodeIDs,
|
||||
ipPreference,
|
||||
options,
|
||||
ping,
|
||||
)...)
|
||||
}
|
||||
items = append(items, p.probePublicTargetCandidates(
|
||||
"exit",
|
||||
0,
|
||||
outNodes,
|
||||
selectedNodeIDs,
|
||||
options,
|
||||
probeTarget,
|
||||
ping,
|
||||
)...)
|
||||
return items
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeCandidateGroupLinks(
|
||||
fromGroup tunnelQualityCandidateGroup,
|
||||
toGroup tunnelQualityCandidateGroup,
|
||||
hopIndex int,
|
||||
selectedNodeIDs map[string]int64,
|
||||
ipPreference string,
|
||||
options diagnosisExecOptions,
|
||||
ping bestExitProbeFunc,
|
||||
) []TunnelQualityCandidateHop {
|
||||
items := make([]TunnelQualityCandidateHop, 0, len(fromGroup.nodes)*len(toGroup.nodes))
|
||||
for _, source := range fromGroup.nodes {
|
||||
for _, target := range toGroup.nodes {
|
||||
item := TunnelQualityCandidateHop{
|
||||
TunnelQualityHop: TunnelQualityHop{
|
||||
FromNodeID: source.NodeID,
|
||||
FromNodeName: source.NodeName,
|
||||
ToNodeID: target.NodeID,
|
||||
ToNodeName: target.NodeName,
|
||||
Latency: -1,
|
||||
Loss: 100,
|
||||
},
|
||||
FromRole: fromGroup.role,
|
||||
ToRole: toGroup.role,
|
||||
HopIndex: hopIndex,
|
||||
Selected: selectedNodeIDs[tunnelQualityGroupKey(fromGroup.role, fromGroup.roleIndex)] == source.NodeID &&
|
||||
selectedNodeIDs[tunnelQualityGroupKey(toGroup.role, toGroup.roleIndex)] == target.NodeID,
|
||||
}
|
||||
|
||||
sourceNode, sourceErr := p.handler.getNodeRecord(source.NodeID)
|
||||
if sourceErr != nil || !isTunnelProbeNodeOnline(sourceNode) {
|
||||
item.ErrorMessage = "来源节点不在线"
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
targetNode, targetErr := p.handler.getNodeRecord(target.NodeID)
|
||||
if targetErr != nil || !isTunnelProbeNodeOnline(targetNode) {
|
||||
item.ErrorMessage = "目标节点不在线"
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
|
||||
targetIP, targetPort, resolveErr := resolveChainProbeTarget(sourceNode, targetNode, target.Port, ipPreference, target.ConnectIP)
|
||||
if resolveErr != nil {
|
||||
item.ErrorMessage = resolveErr.Error()
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
item.TargetIP = targetIP
|
||||
item.TargetPort = targetPort
|
||||
latency, loss, probeErr := ping(source.NodeID, targetIP, targetPort, options)
|
||||
if probeErr != nil {
|
||||
item.ErrorMessage = probeErr.Error()
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
item.Latency = latency
|
||||
item.Loss = loss
|
||||
items = append(items, item)
|
||||
}
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probePublicTargetCandidates(
|
||||
fromRole string,
|
||||
fromIndex int,
|
||||
nodes []chainNodeRecord,
|
||||
selectedNodeIDs map[string]int64,
|
||||
options diagnosisExecOptions,
|
||||
probeTarget tunnelProbeTarget,
|
||||
ping bestExitProbeFunc,
|
||||
) []TunnelQualityCandidateHop {
|
||||
items := make([]TunnelQualityCandidateHop, 0, len(nodes))
|
||||
for _, source := range nodes {
|
||||
item := TunnelQualityCandidateHop{
|
||||
TunnelQualityHop: TunnelQualityHop{
|
||||
FromNodeID: source.NodeID,
|
||||
FromNodeName: source.NodeName,
|
||||
ToNodeName: formatTunnelProbeTarget(probeTarget),
|
||||
Latency: -1,
|
||||
Loss: 100,
|
||||
TargetIP: probeTarget.Host,
|
||||
TargetPort: probeTarget.Port,
|
||||
},
|
||||
FromRole: fromRole,
|
||||
ToRole: "target",
|
||||
HopIndex: fromIndex,
|
||||
Selected: selectedNodeIDs[tunnelQualityGroupKey(fromRole, fromIndex)] == source.NodeID,
|
||||
}
|
||||
sourceNode, sourceErr := p.handler.getNodeRecord(source.NodeID)
|
||||
if sourceErr != nil || !isTunnelProbeNodeOnline(sourceNode) {
|
||||
item.ErrorMessage = "来源节点不在线"
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
latency, loss, probeErr := ping(source.NodeID, probeTarget.Host, probeTarget.Port, options)
|
||||
if probeErr != nil {
|
||||
item.ErrorMessage = probeErr.Error()
|
||||
items = append(items, item)
|
||||
continue
|
||||
}
|
||||
item.Latency = latency
|
||||
item.Loss = loss
|
||||
items = append(items, item)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func isTunnelProbeNodeOnline(node *nodeRecord) bool {
|
||||
return node != nil && (node.IsRemote == 1 || node.Status == 1)
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) firstOnlineChainNode(nodes []chainNodeRecord) (chainNodeRecord, *nodeRecord, bool) {
|
||||
if p == nil || p.handler == nil {
|
||||
return chainNodeRecord{}, nil, false
|
||||
}
|
||||
for _, candidate := range nodes {
|
||||
node, err := p.handler.getNodeRecord(candidate.NodeID)
|
||||
if err == nil && isTunnelProbeNodeOnline(node) {
|
||||
return candidate, node, true
|
||||
}
|
||||
}
|
||||
return chainNodeRecord{}, nil, false
|
||||
}
|
||||
|
||||
func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chainNodeRecord, chainHops [][]chainNodeRecord, outNodes []chainNodeRecord, ipPreference string, options diagnosisExecOptions, probeTarget tunnelProbeTarget, roundPinger bestExitProbeFunc) {
|
||||
if p == nil || p.handler == nil || p.handler.bestExit == nil || len(outNodes) <= 1 {
|
||||
return
|
||||
}
|
||||
@@ -404,9 +686,6 @@ func (p *tunnelQualityProber) probeBestExitOwners(tunnelID int64, inNodes []chai
|
||||
nodeMap[exit.NodeID] = node
|
||||
}
|
||||
}
|
||||
// This best-exit decision cache is per decision round; the display-oriented
|
||||
// tunnel quality snapshot may still collect its own first-exit public probe.
|
||||
roundPinger := newBestExitRoundPinger(p.pingNode)
|
||||
for _, owner := range owners {
|
||||
if nodeMap[owner.NodeID] == nil {
|
||||
continue
|
||||
@@ -444,6 +723,9 @@ func (p *tunnelQualityProber) tcpPingNode(nodeID int64, ip string, port int, opt
|
||||
if nodeErr != nil {
|
||||
return 0, 100, nodeErr
|
||||
}
|
||||
if !isTunnelProbeNodeOnline(node) {
|
||||
return 0, 100, errors.New("节点不在线")
|
||||
}
|
||||
|
||||
var pingData map[string]interface{}
|
||||
var pingErr error
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"slices"
|
||||
"testing"
|
||||
@@ -26,6 +27,9 @@ func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
|
||||
p := newTunnelQualityProber(h)
|
||||
var calls []string
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
if options.pingCount != 1 {
|
||||
t.Fatalf("expected real-time quality probe count 1, got %d", options.pingCount)
|
||||
}
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
@@ -46,6 +50,114 @@ func TestTunnelQualityProberUsesConfiguredProbeTarget(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberSkipsAllOfflineExits(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedQualityForwardTunnel(t, h, 81, []int{0, 0, 0})
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
probeCalls := 0
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
probeCalls++
|
||||
return 0, 100, fmt.Errorf("unexpected probe node=%d target=%s:%d", nodeID, ip, port)
|
||||
}
|
||||
p.probeTunnel(81)
|
||||
|
||||
if probeCalls != 0 {
|
||||
t.Fatalf("expected no TCP probes when all exits are offline, got %d", probeCalls)
|
||||
}
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].Success || snaps[0].ErrorMessage != "出口节点均不在线" {
|
||||
t.Fatalf("expected offline exit snapshot, got %+v", snaps[0])
|
||||
}
|
||||
if snaps[0].EntryToExitLoss != 100 {
|
||||
t.Fatalf("expected 100%% entry-to-exit loss, got %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberUsesOnlineBackupExit(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedQualityForwardTunnel(t, h, 82, []int{0, 1})
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
var calls []string
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
if options.pingCount != 1 {
|
||||
t.Fatalf("expected real-time quality probe count 1, got %d", options.pingCount)
|
||||
}
|
||||
calls = append(calls, fmt.Sprintf("%d|%s|%d", nodeID, ip, port))
|
||||
return 10, 0, nil
|
||||
}
|
||||
p.probeTunnel(82)
|
||||
|
||||
if slices.Contains(calls, "10|10.0.0.30|30030") {
|
||||
t.Fatalf("did not expect probe to offline primary exit, calls=%+v", calls)
|
||||
}
|
||||
if !slices.Contains(calls, "10|10.0.0.31|30031") {
|
||||
t.Fatalf("expected entry probe to online backup exit, calls=%+v", calls)
|
||||
}
|
||||
if !slices.Contains(calls, "31|www.bing.com|443") {
|
||||
t.Fatalf("expected public probe from online backup exit, calls=%+v", calls)
|
||||
}
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 || !snaps[0].Success {
|
||||
t.Fatalf("expected successful backup exit snapshot, got %+v", snaps)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberReportsAllExitCandidateLatencies(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedQualityForwardTunnel(t, h, 83, []int{1, 1})
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
p.probeNode = func(nodeID int64, ip string, port int, options diagnosisExecOptions) (float64, float64, error) {
|
||||
switch fmt.Sprintf("%d|%s|%d", nodeID, ip, port) {
|
||||
case "10|10.0.0.30|30030":
|
||||
return 20, 0, nil
|
||||
case "10|10.0.0.31|30031":
|
||||
return 35, 0, nil
|
||||
case "30|www.bing.com|443":
|
||||
return 50, 0, nil
|
||||
case "31|www.bing.com|443":
|
||||
return 65, 0, nil
|
||||
default:
|
||||
return 0, 100, fmt.Errorf("unexpected probe node=%d target=%s:%d", nodeID, ip, port)
|
||||
}
|
||||
}
|
||||
p.probeTunnel(83)
|
||||
|
||||
snaps := p.GetAll()
|
||||
if len(snaps) != 1 {
|
||||
t.Fatalf("expected one quality snapshot, got %+v", snaps)
|
||||
}
|
||||
if snaps[0].EntryToExitLatency != 20 || snaps[0].ExitToBingLatency != 50 {
|
||||
t.Fatalf("expected primary path metrics to remain unchanged, got %+v", snaps[0])
|
||||
}
|
||||
|
||||
var details tunnelQualityChainDetails
|
||||
if err := json.Unmarshal([]byte(snaps[0].ChainDetails), &details); err != nil {
|
||||
t.Fatalf("decode chain details: %v", err)
|
||||
}
|
||||
assertCandidateHop := func(fromID, toID int64, latency float64, selected bool) {
|
||||
t.Helper()
|
||||
for _, hop := range details.CandidateHops {
|
||||
if hop.FromNodeID == fromID && hop.ToNodeID == toID {
|
||||
if hop.Latency != latency || hop.Selected != selected || hop.ErrorMessage != "" {
|
||||
t.Fatalf("unexpected candidate hop: %+v", hop)
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatalf("candidate hop %d -> %d not found in %+v", fromID, toID, details.CandidateHops)
|
||||
}
|
||||
assertCandidateHop(10, 30, 20, true)
|
||||
assertCandidateHop(10, 31, 35, false)
|
||||
assertCandidateHop(30, 0, 50, true)
|
||||
assertCandidateHop(31, 0, 65, false)
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
seedProbeTargetTunnel(t, h, 78, "quality-target-incomplete", "speed.example.com", 8443)
|
||||
@@ -67,3 +179,69 @@ func TestTunnelQualityProberStoresProbeTargetWhenChainIncomplete(t *testing.T) {
|
||||
t.Fatalf("unexpected snapshot target metadata: %+v", snaps[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberUsesConfiguredInterval(t *testing.T) {
|
||||
h := setupProbeTargetTunnelHandler(t)
|
||||
if err := h.repo.UpsertConfig("monitor_tunnel_quality_interval_sec", "15", time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("upsert interval config: %v", err)
|
||||
}
|
||||
|
||||
p := newTunnelQualityProber(h)
|
||||
if got := p.probeInterval(); got != 15*time.Second {
|
||||
t.Fatalf("probe interval = %s, want 15s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelQualityProberConfigNotificationIsCoalesced(t *testing.T) {
|
||||
p := newTunnelQualityProber(nil)
|
||||
p.NotifyConfigChanged()
|
||||
p.NotifyConfigChanged()
|
||||
|
||||
if got := len(p.wake); got != 1 {
|
||||
t.Fatalf("wake notifications = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelQualityProbeIntervalConfigValue(t *testing.T) {
|
||||
got, err := normalizeAndValidateConfigValue("monitor_tunnel_quality_interval_sec", " 15 ")
|
||||
if err != nil || got != "15" {
|
||||
t.Fatalf("normalize interval = %q, %v", got, err)
|
||||
}
|
||||
if _, err := normalizeAndValidateConfigValue("monitor_tunnel_quality_interval_sec", "0"); err == nil {
|
||||
t.Fatalf("expected invalid interval to be rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func seedQualityForwardTunnel(t *testing.T, h *Handler, tunnelID int64, exitStatuses []int) {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, inx, ip_preference, probe_target_host, probe_target_port)
|
||||
VALUES(?, ?, 1, 2, 'tls', 1, ?, ?, 1, ?, '', '', 0)
|
||||
`, tunnelID, fmt.Sprintf("quality-forward-%d", tunnelID), now, now, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert forwarding tunnel: %v", err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, '1', 10, 30001, 'fifo', 1, 'tls')
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert entry chain: %v", err)
|
||||
}
|
||||
for i, status := range exitStatuses {
|
||||
nodeID := int64(30 + i)
|
||||
port := 30030 + i
|
||||
ip := fmt.Sprintf("10.0.0.%d", nodeID)
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, '', '30000-30100', '', 'v1', 1, 1, 1, ?, ?, ?, '[::]', '[::]', 0)
|
||||
`, nodeID, fmt.Sprintf("exit-%d", i+1), fmt.Sprintf("exit-secret-%d", i+1), ip, ip, now, now, status).Error; err != nil {
|
||||
t.Fatalf("insert exit node %d: %v", nodeID, err)
|
||||
}
|
||||
if err := h.repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, '3', ?, ?, 'fifo', ?, 'tls')
|
||||
`, tunnelID, nodeID, port, i+1).Error; err != nil {
|
||||
t.Fatalf("insert exit chain %d: %v", nodeID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -182,6 +182,8 @@ func requiresAdmin(path string) bool {
|
||||
return true
|
||||
case "/api/v1/config/update", "/api/v1/config/update-single":
|
||||
return true
|
||||
case "/api/v1/license/activate":
|
||||
return true
|
||||
case "/api/v1/announcement/update":
|
||||
return true
|
||||
default:
|
||||
|
||||
@@ -197,6 +197,39 @@ func TestShouldSkipBypassesPublicConfigGet(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLicenseActivateRequiresAdmin(t *testing.T) {
|
||||
if !requiresAdmin("/api/v1/license/activate") {
|
||||
t.Fatal("expected license activation to require admin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJWTRejectsNonAdminLicenseActivation(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(2, "regular_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate token: %v", err)
|
||||
}
|
||||
claims, err := auth.ParseClaims(token, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("parse claims: %v", err)
|
||||
}
|
||||
|
||||
wrapped := JWT(AuthOptions{
|
||||
JWTSecret: secret,
|
||||
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
|
||||
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: claims.IatMs - 1}, nil
|
||||
},
|
||||
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.OK("pass"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", nil)
|
||||
req.Header.Set("Authorization", token)
|
||||
res := httptest.NewRecorder()
|
||||
wrapped.ServeHTTP(res, req)
|
||||
assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作")
|
||||
}
|
||||
|
||||
func TestJWTExpiresAfterSevenDays(t *testing.T) {
|
||||
secret := "unit-test-secret"
|
||||
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
|
||||
@@ -6,24 +6,53 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var AccountID string
|
||||
|
||||
type KeygenClient struct {
|
||||
AccountID string
|
||||
Token string
|
||||
BaseURL string
|
||||
HTTPClient *http.Client
|
||||
}
|
||||
|
||||
type APIError struct {
|
||||
Operation string
|
||||
StatusCode int
|
||||
Body string
|
||||
}
|
||||
|
||||
func (e *APIError) Error() string {
|
||||
return fmt.Sprintf("keygen %s failed: status %d, response: %s", e.Operation, e.StatusCode, e.Body)
|
||||
}
|
||||
|
||||
func (e *APIError) HasCode(code string) bool {
|
||||
return e != nil && hasKeygenErrorCode([]byte(e.Body), code)
|
||||
}
|
||||
|
||||
const defaultAPIBaseURL = "https://api.keygen.sh/v1"
|
||||
|
||||
func NewKeygenClient(accountID, token string) *KeygenClient {
|
||||
return &KeygenClient{
|
||||
AccountID: accountID,
|
||||
Token: token,
|
||||
AccountID: accountID,
|
||||
Token: token,
|
||||
BaseURL: defaultAPIBaseURL,
|
||||
HTTPClient: &http.Client{Timeout: 10 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *KeygenClient) apiURL(path string) string {
|
||||
baseURL := strings.TrimRight(c.BaseURL, "/")
|
||||
if baseURL == "" {
|
||||
baseURL = defaultAPIBaseURL
|
||||
}
|
||||
return fmt.Sprintf("%s/accounts/%s/%s", baseURL, c.AccountID, strings.TrimLeft(path, "/"))
|
||||
}
|
||||
|
||||
type ValidateResponse struct {
|
||||
Meta struct {
|
||||
Valid bool `json:"valid"`
|
||||
@@ -35,6 +64,7 @@ type ValidateResponse struct {
|
||||
Expiry string `json:"expiry"`
|
||||
} `json:"attributes"`
|
||||
} `json:"data"`
|
||||
MachineID string `json:"-"`
|
||||
}
|
||||
|
||||
type ActivateMachineRequest struct {
|
||||
@@ -54,8 +84,37 @@ type ActivateMachineRequest struct {
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
type keygenErrorResponse struct {
|
||||
Errors []struct {
|
||||
Code string `json:"code"`
|
||||
} `json:"errors"`
|
||||
}
|
||||
|
||||
type MachineResponse struct {
|
||||
Data struct {
|
||||
ID string `json:"id"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
func hasKeygenErrorCode(body []byte, code string) bool {
|
||||
var resp keygenErrorResponse
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
return false
|
||||
}
|
||||
for _, item := range resp.Errors {
|
||||
if strings.EqualFold(strings.TrimSpace(item.Code), code) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string) (*ValidateResponse, error) {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
|
||||
return c.ValidateKeyWithMachine(key, fingerprint, "")
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ValidateKeyWithMachine(key, fingerprint, machineID string) (*ValidateResponse, error) {
|
||||
url := c.apiURL("licenses/actions/validate-key")
|
||||
|
||||
meta := map[string]interface{}{
|
||||
"key": key,
|
||||
@@ -66,6 +125,14 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
|
||||
"fingerprint": fingerprint,
|
||||
}
|
||||
}
|
||||
if machineID != "" {
|
||||
scope, _ := meta["scope"].(map[string]interface{})
|
||||
if scope == nil {
|
||||
scope = make(map[string]interface{})
|
||||
meta["scope"] = scope
|
||||
}
|
||||
scope["machine"] = machineID
|
||||
}
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"meta": meta,
|
||||
@@ -91,7 +158,8 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
var valResp ValidateResponse
|
||||
@@ -102,8 +170,44 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
|
||||
return &valResp, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) GetMachineID(fingerprint string) (string, error) {
|
||||
machineURL := c.apiURL("machines/" + url.PathEscape(fingerprint))
|
||||
req, err := http.NewRequest(http.MethodGet, machineURL, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Accept", "application/vnd.api+json")
|
||||
if c.Token != "" {
|
||||
if !strings.HasPrefix(c.Token, "Bearer ") && !strings.HasPrefix(c.Token, "License ") {
|
||||
req.Header.Set("Authorization", "License "+c.Token)
|
||||
} else {
|
||||
req.Header.Set("Authorization", c.Token)
|
||||
}
|
||||
}
|
||||
|
||||
resp, err := c.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return "", &APIError{Operation: "retrieve machine", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
var machineResp MachineResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&machineResp); err != nil {
|
||||
return "", err
|
||||
}
|
||||
machineID := strings.TrimSpace(machineResp.Data.ID)
|
||||
if machineID == "" {
|
||||
return "", fmt.Errorf("failed to retrieve machine: empty machine id")
|
||||
}
|
||||
return machineID, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
|
||||
url := c.apiURL("licenses/actions/validate-key")
|
||||
|
||||
reqBody := map[string]interface{}{
|
||||
"meta": map[string]string{
|
||||
@@ -130,7 +234,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
var valResp ValidateResponse
|
||||
@@ -141,8 +246,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
|
||||
return &valResp, nil
|
||||
}
|
||||
|
||||
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
|
||||
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/machines", c.AccountID)
|
||||
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) (string, error) {
|
||||
url := c.apiURL("machines")
|
||||
|
||||
var reqBody ActivateMachineRequest
|
||||
reqBody.Data.Type = "machines"
|
||||
@@ -165,23 +270,24 @@ func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
|
||||
|
||||
resp, err := c.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
|
||||
if resp.StatusCode == http.StatusCreated || resp.StatusCode == http.StatusOK {
|
||||
return nil
|
||||
}
|
||||
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
|
||||
if resp.StatusCode == http.StatusConflict || resp.StatusCode == http.StatusUnprocessableEntity {
|
||||
if strings.Contains(string(body), "FINGERPRINT_TAKEN") || strings.Contains(string(body), "MACHINE_LIMIT_EXCEEDED") {
|
||||
// Machine already registered to this license or limit reached because it's already us.
|
||||
// The subsequent ValidateKey check will determine if the existing machine is actually us.
|
||||
return nil
|
||||
var machineResp MachineResponse
|
||||
if json.Unmarshal(body, &machineResp) == nil {
|
||||
return strings.TrimSpace(machineResp.Data.ID), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to activate machine: status %d, response: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
if resp.StatusCode == http.StatusUnprocessableEntity && hasKeygenErrorCode(body, "FINGERPRINT_TAKEN") {
|
||||
// Machine activation is idempotent. Keygen scopes fingerprint uniqueness
|
||||
// to the target license, so this means the same machine is already bound.
|
||||
return "", nil
|
||||
}
|
||||
|
||||
return "", &APIError{Operation: "activate machine", StatusCode: resp.StatusCode, Body: string(body)}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
package license
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateKeyWithMachineSendsFingerprintAndMachineScopes(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var body struct {
|
||||
Meta struct {
|
||||
Scope map[string]string `json:"scope"`
|
||||
} `json:"meta"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
if body.Meta.Scope["fingerprint"] != "fingerprint" {
|
||||
t.Fatalf("fingerprint scope = %q", body.Meta.Scope["fingerprint"])
|
||||
}
|
||||
if body.Meta.Scope["machine"] != "machine-id" {
|
||||
t.Fatalf("machine scope = %q", body.Meta.Scope["machine"])
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{}}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "")
|
||||
client.BaseURL = server.URL
|
||||
validation, err := client.ValidateKeyWithMachine("license-key", "fingerprint", "machine-id")
|
||||
if err != nil || !validation.Meta.Valid {
|
||||
t.Fatalf("ValidateKeyWithMachine() validation=%+v err=%v", validation, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetMachineIDRetrievesMachineByFingerprint(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet || !strings.HasSuffix(r.URL.Path, "/machines/fingerprint") {
|
||||
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
|
||||
}
|
||||
if r.Header.Get("Authorization") != "License license-key" {
|
||||
t.Fatalf("authorization = %q", r.Header.Get("Authorization"))
|
||||
}
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
machineID, err := client.GetMachineID("fingerprint")
|
||||
if err != nil || machineID != "machine-id" {
|
||||
t.Fatalf("GetMachineID() = %q, %v", machineID, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateMachineTreatsFingerprintTakenAsIdempotent(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
|
||||
if machineID, err := client.ActivateMachine("license-id", "fingerprint"); err != nil || machineID != "" {
|
||||
t.Fatalf("ActivateMachine() error = %v, want idempotent success", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateMachineReturnsCreatedMachineID(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
machineID, err := client.ActivateMachine("license-id", "fingerprint")
|
||||
if err != nil || machineID != "machine-id" {
|
||||
t.Fatalf("ActivateMachine() = %q, %v", machineID, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateMachineRejectsMachineLimitWithoutFingerprintTaken(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnprocessableEntity)
|
||||
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewKeygenClient("account-id", "license-key")
|
||||
client.BaseURL = server.URL
|
||||
|
||||
_, err := client.ActivateMachine("license-id", "fingerprint")
|
||||
if err == nil || !strings.Contains(err.Error(), "MACHINE_LIMIT_EXCEEDED") {
|
||||
t.Fatalf("ActivateMachine() error = %v, want machine limit failure", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package monitoring
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
ConfigTunnelQualityProbeIntervalSec = "monitor_tunnel_quality_interval_sec"
|
||||
DefaultTunnelQualityProbeIntervalSec = 1
|
||||
MinTunnelQualityProbeIntervalSec = 1
|
||||
MaxTunnelQualityProbeIntervalSec = 3600
|
||||
)
|
||||
|
||||
func TunnelQualityProbeIntervalSecondsFromConfigMap(cfg map[string]string) int {
|
||||
if cfg == nil {
|
||||
return DefaultTunnelQualityProbeIntervalSec
|
||||
}
|
||||
seconds, err := parseTunnelQualityProbeIntervalSeconds(cfg[ConfigTunnelQualityProbeIntervalSec])
|
||||
if err != nil {
|
||||
return DefaultTunnelQualityProbeIntervalSec
|
||||
}
|
||||
return seconds
|
||||
}
|
||||
|
||||
func NormalizeTunnelQualityProbeIntervalSeconds(value string) (string, error) {
|
||||
seconds, err := parseTunnelQualityProbeIntervalSeconds(value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strconv.Itoa(seconds), nil
|
||||
}
|
||||
|
||||
func parseTunnelQualityProbeIntervalSeconds(value string) (int, error) {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
if trimmed == "" {
|
||||
return 0, fmt.Errorf("隧道质量探测间隔不能为空")
|
||||
}
|
||||
seconds, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("隧道质量探测间隔必须是整数")
|
||||
}
|
||||
if seconds < MinTunnelQualityProbeIntervalSec || seconds > MaxTunnelQualityProbeIntervalSec {
|
||||
return 0, fmt.Errorf(
|
||||
"隧道质量探测间隔必须在 %d 到 %d 秒之间",
|
||||
MinTunnelQualityProbeIntervalSec,
|
||||
MaxTunnelQualityProbeIntervalSec,
|
||||
)
|
||||
}
|
||||
return seconds, nil
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package monitoring
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestTunnelQualityProbeIntervalSecondsFromConfigMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg map[string]string
|
||||
want int
|
||||
}{
|
||||
{name: "missing config", cfg: nil, want: DefaultTunnelQualityProbeIntervalSec},
|
||||
{name: "configured", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "15"}, want: 15},
|
||||
{name: "invalid", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "0"}, want: DefaultTunnelQualityProbeIntervalSec},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := TunnelQualityProbeIntervalSecondsFromConfigMap(tt.cfg); got != tt.want {
|
||||
t.Fatalf("interval = %d, want %d", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeTunnelQualityProbeIntervalSeconds(t *testing.T) {
|
||||
for _, value := range []string{"1", "15", "3600"} {
|
||||
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err != nil || got != value {
|
||||
t.Fatalf("normalize %q = %q, %v", value, got, err)
|
||||
}
|
||||
}
|
||||
|
||||
for _, value := range []string{"", "0", "3601", "1.5", "abc"} {
|
||||
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err == nil {
|
||||
t.Fatalf("normalize %q unexpectedly succeeded with %q", value, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type CounterSample struct {
|
||||
ForwardID int64
|
||||
Direction string
|
||||
Protocol string
|
||||
Bytes uint64
|
||||
Packets uint64
|
||||
}
|
||||
|
||||
func ParseCounterComment(comment string) (CounterSample, bool) {
|
||||
parts := strings.Split(comment, " ")
|
||||
if len(parts) != 4 || parts[0] != "flvx" {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
if !strings.HasPrefix(parts[1], "forward:") {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
forwardText := strings.TrimPrefix(parts[1], "forward:")
|
||||
forwardID, err := strconv.ParseInt(forwardText, 10, 64)
|
||||
if err != nil || forwardID <= 0 {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
direction := parts[2]
|
||||
if direction != CounterDirectionToTarget && direction != CounterDirectionFromTarget {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
protocol := parts[3]
|
||||
if protocol != "tcp" && protocol != "udp" {
|
||||
return CounterSample{}, false
|
||||
}
|
||||
return CounterSample{
|
||||
ForwardID: forwardID,
|
||||
Direction: direction,
|
||||
Protocol: protocol,
|
||||
}, true
|
||||
}
|
||||
|
||||
func ParseCounterSamples(raw []byte) ([]CounterSample, error) {
|
||||
var doc nftListTable
|
||||
if err := json.Unmarshal(raw, &doc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
samples := make([]CounterSample, 0)
|
||||
for _, item := range doc.Nftables {
|
||||
ruleRaw, ok := item["rule"]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
var rule nftCounterRule
|
||||
if err := json.Unmarshal(ruleRaw, &rule); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if rule.Table != "flvx" || rule.Chain != "forward" {
|
||||
continue
|
||||
}
|
||||
|
||||
sample, ok, err := parseCounterRule(rule)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
samples = append(samples, sample)
|
||||
}
|
||||
return samples, nil
|
||||
}
|
||||
|
||||
type nftListTable struct {
|
||||
Nftables []map[string]json.RawMessage `json:"nftables"`
|
||||
}
|
||||
|
||||
type nftCounterRule struct {
|
||||
Table string `json:"table"`
|
||||
Chain string `json:"chain"`
|
||||
Comment string `json:"comment"`
|
||||
Expr []map[string]json.RawMessage `json:"expr"`
|
||||
}
|
||||
|
||||
type nftCounter struct {
|
||||
Bytes uint64 `json:"bytes"`
|
||||
Packets uint64 `json:"packets"`
|
||||
}
|
||||
|
||||
func parseCounterRule(rule nftCounterRule) (CounterSample, bool, error) {
|
||||
var (
|
||||
counter nftCounter
|
||||
hasCounter bool
|
||||
comment = rule.Comment
|
||||
)
|
||||
|
||||
for _, expr := range rule.Expr {
|
||||
if rawCounter, ok := expr["counter"]; ok {
|
||||
if err := json.Unmarshal(rawCounter, &counter); err != nil {
|
||||
return CounterSample{}, false, err
|
||||
}
|
||||
hasCounter = true
|
||||
continue
|
||||
}
|
||||
if rawComment, ok := expr["comment"]; ok && strings.TrimSpace(comment) == "" {
|
||||
if err := json.Unmarshal(rawComment, &comment); err != nil {
|
||||
return CounterSample{}, false, err
|
||||
}
|
||||
}
|
||||
}
|
||||
if !hasCounter {
|
||||
return CounterSample{}, false, nil
|
||||
}
|
||||
|
||||
sample, ok := ParseCounterComment(comment)
|
||||
if !ok {
|
||||
return CounterSample{}, false, nil
|
||||
}
|
||||
sample.Bytes = counter.Bytes
|
||||
sample.Packets = counter.Packets
|
||||
return sample, true, nil
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
package nftables
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseCounterCommentAcceptsValidToTargetTCP(t *testing.T) {
|
||||
sample, ok := ParseCounterComment("flvx forward:42 to-target tcp")
|
||||
if !ok {
|
||||
t.Fatal("expected comment to parse")
|
||||
}
|
||||
if sample.ForwardID != 42 ||
|
||||
sample.Direction != CounterDirectionToTarget ||
|
||||
sample.Protocol != "tcp" {
|
||||
t.Fatalf("unexpected sample: %+v", sample)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterCommentRejectsDNAT(t *testing.T) {
|
||||
if sample, ok := ParseCounterComment("flvx forward:42 dnat tcp"); ok {
|
||||
t.Fatalf("expected dnat comment to be rejected, got %+v", sample)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesParsesForwardBillableCounters(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"metainfo": {"json_schema_version": 1}},
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"handle": 10,
|
||||
"comment": "flvx forward:42 to-target tcp",
|
||||
"expr": [
|
||||
{"match": {"left": {"payload": {"protocol": "ip", "field": "daddr"}}, "op": "==", "right": "198.51.100.20"}},
|
||||
{"counter": {"packets": 7, "bytes": 4096}}
|
||||
]
|
||||
}},
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"handle": 11,
|
||||
"comment": "flvx forward:42 from-target udp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 9, "bytes": 8192}}
|
||||
]
|
||||
}},
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "prerouting",
|
||||
"handle": 12,
|
||||
"comment": "flvx forward:42 dnat tcp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 100, "bytes": 65536}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
|
||||
samples, err := ParseCounterSamples(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCounterSamples: %v", err)
|
||||
}
|
||||
if len(samples) != 2 {
|
||||
t.Fatalf("expected 2 samples, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
|
||||
want := []CounterSample{
|
||||
{ForwardID: 42, Direction: CounterDirectionToTarget, Protocol: "tcp", Bytes: 4096, Packets: 7},
|
||||
{ForwardID: 42, Direction: CounterDirectionFromTarget, Protocol: "udp", Bytes: 8192, Packets: 9},
|
||||
}
|
||||
for i := range want {
|
||||
if samples[i] != want[i] {
|
||||
t.Fatalf("sample %d: expected %+v, got %+v", i, want[i], samples[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesUsesRuleLevelComment(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"comment": "flvx forward:77 to-target udp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 3, "bytes": 2048}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
|
||||
samples, err := ParseCounterSamples(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCounterSamples: %v", err)
|
||||
}
|
||||
if len(samples) != 1 {
|
||||
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
want := CounterSample{
|
||||
ForwardID: 77,
|
||||
Direction: CounterDirectionToTarget,
|
||||
Protocol: "udp",
|
||||
Bytes: 2048,
|
||||
Packets: 3,
|
||||
}
|
||||
if samples[0] != want {
|
||||
t.Fatalf("expected %+v, got %+v", want, samples[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesUsesExprLevelComment(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"expr": [
|
||||
{"counter": {"packets": 4, "bytes": 3072}},
|
||||
{"comment": "flvx forward:78 from-target tcp"}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
|
||||
samples, err := ParseCounterSamples(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseCounterSamples: %v", err)
|
||||
}
|
||||
if len(samples) != 1 {
|
||||
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
want := CounterSample{
|
||||
ForwardID: 78,
|
||||
Direction: CounterDirectionFromTarget,
|
||||
Protocol: "tcp",
|
||||
Bytes: 3072,
|
||||
Packets: 4,
|
||||
}
|
||||
if samples[0] != want {
|
||||
t.Fatalf("expected %+v, got %+v", want, samples[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesMalformedJSONReturnsError(t *testing.T) {
|
||||
if _, err := ParseCounterSamples([]byte(`{"nftables": [`)); err == nil {
|
||||
t.Fatal("expected malformed JSON error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCounterSamplesMalformedRuleJSONReturnsError(t *testing.T) {
|
||||
raw := []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"comment": "flvx forward:42 to-target tcp",
|
||||
"expr": [
|
||||
{"counter": {"packets": "bad", "bytes": 4096}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)
|
||||
if _, err := ParseCounterSamples(raw); err == nil {
|
||||
t.Fatal("expected malformed rule JSON error")
|
||||
}
|
||||
}
|
||||
@@ -46,6 +46,17 @@ func (m *Manager) Clear(ctx context.Context, cfg SSHConfig) error {
|
||||
return m.runner.ApplyScript(ctx, cfg, script)
|
||||
}
|
||||
|
||||
func (m *Manager) CollectCounters(ctx context.Context, cfg SSHConfig) ([]CounterSample, error) {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
raw, err := m.runner.ListTableJSON(ctx, cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return ParseCounterSamples(raw)
|
||||
}
|
||||
|
||||
func (m *Manager) ensureInitialized() error {
|
||||
if m == nil || m.runner == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
|
||||
@@ -8,9 +8,11 @@ import (
|
||||
)
|
||||
|
||||
type fakeRunner struct {
|
||||
scripts []string
|
||||
err error
|
||||
testErr error
|
||||
scripts []string
|
||||
err error
|
||||
testErr error
|
||||
listJSON []byte
|
||||
listJSONErr error
|
||||
}
|
||||
|
||||
func (f *fakeRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
@@ -22,12 +24,16 @@ func (f *fakeRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
|
||||
return f.listJSON, f.listJSONErr
|
||||
}
|
||||
|
||||
func TestManagerReconcileAppliesRenderedScript(t *testing.T) {
|
||||
runner := &fakeRunner{}
|
||||
manager := NewManager(runner)
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}},
|
||||
Rules: []Rule{{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}},
|
||||
}
|
||||
|
||||
result, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, plan)
|
||||
@@ -37,7 +43,7 @@ func TestManagerReconcileAppliesRenderedScript(t *testing.T) {
|
||||
if len(runner.scripts) != 1 {
|
||||
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
|
||||
}
|
||||
if !strings.Contains(runner.scripts[0], "flvx forward:42 tcp") {
|
||||
if !strings.Contains(runner.scripts[0], `flvx forward:42 dnat tcp`) {
|
||||
t.Fatalf("script missing forward comment:\n%s", runner.scripts[0])
|
||||
}
|
||||
if result.NodeID != 7 || result.Hashes[42] == "" {
|
||||
@@ -80,6 +86,41 @@ func TestManagerTestPassesThroughRunnerError(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerCollectCountersParsesRunnerTableJSON(t *testing.T) {
|
||||
runner := &fakeRunner{listJSON: []byte(`{
|
||||
"nftables": [
|
||||
{"rule": {
|
||||
"family": "inet",
|
||||
"table": "flvx",
|
||||
"chain": "forward",
|
||||
"comment": "flvx forward:77 to-target tcp",
|
||||
"expr": [
|
||||
{"counter": {"packets": 3, "bytes": 2048}}
|
||||
]
|
||||
}}
|
||||
]
|
||||
}`)}
|
||||
manager := NewManager(runner)
|
||||
|
||||
samples, err := manager.CollectCounters(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"})
|
||||
if err != nil {
|
||||
t.Fatalf("CollectCounters: %v", err)
|
||||
}
|
||||
if len(samples) != 1 {
|
||||
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
|
||||
}
|
||||
want := CounterSample{
|
||||
ForwardID: 77,
|
||||
Direction: CounterDirectionToTarget,
|
||||
Protocol: "tcp",
|
||||
Bytes: 2048,
|
||||
Packets: 3,
|
||||
}
|
||||
if samples[0] != want {
|
||||
t.Fatalf("expected %+v, got %+v", want, samples[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
|
||||
cfg := SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}
|
||||
plan := NodePlan{NodeID: 7}
|
||||
@@ -98,6 +139,10 @@ func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
|
||||
t.Fatalf("expected not initialized error from nil manager Clear, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := nilManager.CollectCounters(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager CollectCounters, got %v", err)
|
||||
}
|
||||
|
||||
manager := &Manager{}
|
||||
if err := manager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Test, got %v", err)
|
||||
@@ -110,4 +155,8 @@ func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
|
||||
if err := manager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Clear, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := manager.CollectCounters(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from CollectCounters, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,14 +15,18 @@ func RenderTable(plan NodePlan) string {
|
||||
b.WriteString(" chain prerouting {\n")
|
||||
b.WriteString(" type nat hook prerouting priority dstnat; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
family := nftAddressFamily(rule.TargetHost)
|
||||
dnatFamily := ""
|
||||
if family != "" {
|
||||
dnatFamily = family + " "
|
||||
}
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" %s dport %d dnat %s to %s comment \"flvx forward:%d %s\"\n",
|
||||
b.WriteString(fmt.Sprintf(" %s dport %d counter dnat %sto %s comment %q\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
dnatFamilyPrefix(rule.TargetHost),
|
||||
dnatFamily,
|
||||
formatDNATTarget(rule.TargetHost, rule.TargetPort),
|
||||
rule.ForwardID,
|
||||
protocol,
|
||||
counterComment(rule.ForwardID, CounterDirectionDNAT, protocol),
|
||||
))
|
||||
}
|
||||
}
|
||||
@@ -35,11 +39,42 @@ func RenderTable(plan NodePlan) string {
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain forward {\n")
|
||||
b.WriteString(" type filter hook forward priority filter; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
family := nftAddressFamily(rule.TargetHost)
|
||||
if family == "" {
|
||||
continue
|
||||
}
|
||||
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
|
||||
))
|
||||
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
family,
|
||||
targetHost,
|
||||
protocol,
|
||||
rule.TargetPort,
|
||||
counterComment(rule.ForwardID, CounterDirectionFromTarget, protocol),
|
||||
))
|
||||
}
|
||||
}
|
||||
b.WriteString(" }\n")
|
||||
b.WriteString("}\n")
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func counterComment(forwardID int64, direction, protocol string) string {
|
||||
return fmt.Sprintf("flvx forward:%d %s %s", forwardID, direction, protocol)
|
||||
}
|
||||
|
||||
func RuleHash(rule Rule) string {
|
||||
protocols := normalizedProtocols(rule.Protocols)
|
||||
sum := sha256.Sum256([]byte(fmt.Sprintf("%d|%d|%s|%d|%s",
|
||||
@@ -100,14 +135,14 @@ func formatDNATTarget(host string, port int) string {
|
||||
return fmt.Sprintf("%s:%d", trimmed, port)
|
||||
}
|
||||
|
||||
func dnatFamilyPrefix(host string) string {
|
||||
func nftAddressFamily(host string) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
ip := net.ParseIP(trimmed)
|
||||
if ip == nil {
|
||||
return ""
|
||||
}
|
||||
if ip.To4() != nil {
|
||||
return "ip"
|
||||
if ip.To4() == nil {
|
||||
return "ip6"
|
||||
}
|
||||
return "ip6"
|
||||
return "ip"
|
||||
}
|
||||
|
||||
@@ -23,8 +23,8 @@ func TestRenderTableIncludesDNATAndMasquerade(t *testing.T) {
|
||||
"table inet flvx",
|
||||
"type nat hook prerouting priority dstnat; policy accept;",
|
||||
"type nat hook postrouting priority srcnat; policy accept;",
|
||||
"tcp dport 24000 dnat ip to 198.51.100.20:443 comment \"flvx forward:42 tcp\"",
|
||||
"udp dport 24000 dnat ip to 198.51.100.20:443 comment \"flvx forward:42 udp\"",
|
||||
"tcp dport 24000 counter dnat ip to 198.51.100.20:443 comment \"flvx forward:42 dnat tcp\"",
|
||||
"udp dport 24000 counter dnat ip to 198.51.100.20:443 comment \"flvx forward:42 dnat udp\"",
|
||||
"masquerade comment \"flvx masquerade\"",
|
||||
}
|
||||
for _, part := range expectedParts {
|
||||
@@ -46,6 +46,109 @@ func TestRenderTableBracketsIPv6Target(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 42,
|
||||
InPort: 12345,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
|
||||
`udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`,
|
||||
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`meta l4proto tcp ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
|
||||
`meta l4proto udp ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
|
||||
`meta l4proto udp ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 43,
|
||||
InPort: 12346,
|
||||
TargetHost: "2001:db8::20",
|
||||
TargetPort: 8443,
|
||||
Protocols: []string{"tcp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
|
||||
`meta l4proto tcp ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
`meta l4proto tcp ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{
|
||||
{ForwardID: 42, InPort: 12345, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
{ForwardID: 43, InPort: 12346, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
wantLines := []string{
|
||||
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
|
||||
`meta l4proto tcp ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
|
||||
}
|
||||
for _, want := range wantLines {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTablePreservesHostnameDNATAndSkipsAccountingCounters(t *testing.T) {
|
||||
plan := NodePlan{
|
||||
NodeID: 7,
|
||||
Rules: []Rule{{
|
||||
ForwardID: 44,
|
||||
InPort: 12347,
|
||||
TargetHost: "example.com",
|
||||
TargetPort: 9443,
|
||||
Protocols: []string{"tcp"},
|
||||
}},
|
||||
}
|
||||
|
||||
got := RenderTable(plan)
|
||||
want := `tcp dport 12347 counter dnat to example.com:9443 comment "flvx forward:44 dnat tcp"`
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("RenderTable() missing %q\n%s", want, got)
|
||||
}
|
||||
unwantedLines := []string{
|
||||
`dnat ip to example.com`,
|
||||
`ip daddr example.com`,
|
||||
`ip saddr example.com`,
|
||||
}
|
||||
for _, unwanted := range unwantedLines {
|
||||
if strings.Contains(got, unwanted) {
|
||||
t.Fatalf("RenderTable() unexpectedly contains %q\n%s", unwanted, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIsStable(t *testing.T) {
|
||||
rule := Rule{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}
|
||||
if RuleHash(rule) != RuleHash(rule) {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
type Runner interface {
|
||||
ApplyScript(ctx context.Context, cfg SSHConfig, script string) error
|
||||
Test(ctx context.Context, cfg SSHConfig) error
|
||||
ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error)
|
||||
}
|
||||
|
||||
type SSHRunner struct {
|
||||
@@ -25,26 +26,70 @@ func NewSSHRunner() *SSHRunner {
|
||||
}
|
||||
|
||||
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
|
||||
nft := nftBinary(cfg)
|
||||
tableName := fmt.Sprintf("flvx_capability_%d", time.Now().UnixNano())
|
||||
return r.run(ctx, cfg, buildCapabilityCheckCommand(nft, tableName))
|
||||
}
|
||||
|
||||
func buildCapabilityCheckCommand(nft, tableName string) string {
|
||||
script := RenderTable(NodePlan{
|
||||
Rules: []Rule{{
|
||||
ForwardID: 1,
|
||||
InPort: 12345,
|
||||
TargetHost: "192.0.2.1",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}},
|
||||
})
|
||||
script = strings.Replace(script, "table inet flvx {", "table inet "+tableName+" {", 1)
|
||||
return "set -eu\n" +
|
||||
"command -v nft >/dev/null 2>&1\n" +
|
||||
nft + " --version >/dev/null 2>&1\n" +
|
||||
"tmp=$(mktemp /tmp/flvx-nft-capability-XXXXXX.nft)\n" +
|
||||
"trap 'rm -f \"$tmp\"' EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
"if ! " + nft + " -c -f \"$tmp\"; then\n" +
|
||||
" echo 'nftables cannot validate the generated FLVX rules' >&2\n" +
|
||||
" exit 1\n" +
|
||||
"fi"
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
nft := nftBinary(cfg)
|
||||
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
|
||||
"cleanup() {\n" +
|
||||
" rm -f \"$tmp\"\n" +
|
||||
"}\n" +
|
||||
"trap cleanup EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
nft + " -c -f \"$tmp\"\n" +
|
||||
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
|
||||
" " + nft + " delete table inet flvx\n" +
|
||||
"fi\n" +
|
||||
nft + " -f \"$tmp\""
|
||||
command := buildApplyCommand(nftBinary(cfg), script)
|
||||
return r.run(ctx, cfg, command)
|
||||
}
|
||||
|
||||
func buildApplyCommand(nft, script string) string {
|
||||
return "set -eu\n" +
|
||||
"tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft)\n" +
|
||||
"batch=$(mktemp /tmp/flvx-nft-batch-XXXXXX.nft) || { rm -f \"$tmp\"; exit 1; }\n" +
|
||||
"cleanup() {\n" +
|
||||
" rm -f \"$tmp\" \"$batch\"\n" +
|
||||
"}\n" +
|
||||
"trap cleanup EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
|
||||
" { printf '%s\\n' 'delete table inet flvx'; cat \"$tmp\"; } > \"$batch\"\n" +
|
||||
"else\n" +
|
||||
" cp \"$tmp\" \"$batch\"\n" +
|
||||
"fi\n" +
|
||||
"if ! " + nft + " -c -f \"$batch\"; then\n" +
|
||||
" echo 'nftables rule validation failed; active rules were preserved' >&2\n" +
|
||||
" exit 1\n" +
|
||||
"fi\n" +
|
||||
nft + " -f \"$batch\""
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
|
||||
return r.runOutput(ctx, cfg, nftBinary(cfg)+" -j list table inet flvx")
|
||||
}
|
||||
|
||||
func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) error {
|
||||
_, err := r.runOutput(ctx, cfg, command)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *SSHRunner) runOutput(ctx context.Context, cfg SSHConfig, command string) ([]byte, error) {
|
||||
timeout := r.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 15 * time.Second
|
||||
@@ -54,31 +99,33 @@ func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) erro
|
||||
|
||||
clientConfig, err := buildSSHClientConfig(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
addr := net.JoinHostPort(strings.TrimSpace(cfg.Host), fmt.Sprintf("%d", normalizedSSHPort(cfg.Port)))
|
||||
dialer := net.Dialer{Timeout: timeout}
|
||||
conn, err := dialer.DialContext(runCtx, "tcp", addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SSH 连接失败: %w", err)
|
||||
return nil, fmt.Errorf("SSH 连接失败: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SSH 认证失败: %w", err)
|
||||
return nil, fmt.Errorf("SSH 认证失败: %w", err)
|
||||
}
|
||||
client := ssh.NewClient(sshConn, chans, reqs)
|
||||
defer client.Close()
|
||||
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return fmt.Errorf("SSH 会话创建失败: %w", err)
|
||||
return nil, fmt.Errorf("SSH 会话创建失败: %w", err)
|
||||
}
|
||||
defer session.Close()
|
||||
|
||||
var stdout bytes.Buffer
|
||||
var stderr bytes.Buffer
|
||||
session.Stdout = &stdout
|
||||
session.Stderr = &stderr
|
||||
|
||||
done := make(chan error, 1)
|
||||
@@ -89,16 +136,16 @@ func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) erro
|
||||
select {
|
||||
case <-runCtx.Done():
|
||||
_ = session.Close()
|
||||
return fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
|
||||
return nil, fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
message := strings.TrimSpace(stderr.String())
|
||||
if message != "" {
|
||||
return fmt.Errorf("远程执行失败: %s: %w", message, err)
|
||||
return nil, fmt.Errorf("远程执行失败: %s: %w", message, err)
|
||||
}
|
||||
return fmt.Errorf("远程执行失败: %w", err)
|
||||
return nil, fmt.Errorf("远程执行失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
return stdout.Bytes(), nil
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -5,10 +5,113 @@ import (
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestBuildApplyCommandStopsAfterValidationFailure(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
logPath := filepath.Join(dir, "calls.log")
|
||||
applyMarker := filepath.Join(dir, "applied")
|
||||
nftPath := filepath.Join(dir, "nft")
|
||||
fake := `#!/bin/sh
|
||||
echo "$*" >> "` + logPath + `"
|
||||
if [ "$1" = "list" ]; then
|
||||
exit 0
|
||||
fi
|
||||
if [ "$1" = "-c" ]; then
|
||||
exit 1
|
||||
fi
|
||||
touch "` + applyMarker + `"
|
||||
`
|
||||
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
|
||||
t.Fatalf("write fake nft: %v", err)
|
||||
}
|
||||
|
||||
command := buildApplyCommand(nftPath, "table inet flvx { }")
|
||||
result := exec.Command("sh", "-c", command)
|
||||
if err := result.Run(); err == nil {
|
||||
t.Fatal("expected validation failure")
|
||||
}
|
||||
if _, err := os.Stat(applyMarker); !os.IsNotExist(err) {
|
||||
t.Fatalf("apply ran after validation failure, stat err=%v", err)
|
||||
}
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read fake nft calls: %v", err)
|
||||
}
|
||||
if strings.Count(string(calls), "-f ") != 1 {
|
||||
t.Fatalf("expected validation only, got calls:\n%s", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCapabilityCheckCommandValidatesRenderedRulesWithoutApplying(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
logPath := filepath.Join(dir, "calls.log")
|
||||
nftPath := filepath.Join(dir, "nft")
|
||||
fake := `#!/bin/sh
|
||||
echo "$*" >> "` + logPath + `"
|
||||
if [ "$1" = "--version" ] || [ "$1" = "-c" ]; then
|
||||
exit 0
|
||||
fi
|
||||
exit 1
|
||||
`
|
||||
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
|
||||
t.Fatalf("write fake nft: %v", err)
|
||||
}
|
||||
|
||||
command := strings.Replace(buildCapabilityCheckCommand(nftPath, "flvx_capability_test"), "command -v nft", "command -v "+nftPath, 1)
|
||||
result := exec.Command("sh", "-c", command)
|
||||
if output, err := result.CombinedOutput(); err != nil {
|
||||
t.Fatalf("capability command failed: %v: %s", err, output)
|
||||
}
|
||||
calls, err := os.ReadFile(logPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read fake nft calls: %v", err)
|
||||
}
|
||||
if strings.Count(string(calls), "-c -f ") != 1 || strings.Contains(string(calls), "\n-f ") {
|
||||
t.Fatalf("expected one check-only invocation, got calls:\n%s", calls)
|
||||
}
|
||||
if !strings.Contains(command, "table inet flvx_capability_test") || !strings.Contains(command, "meta l4proto tcp ct original proto-dst") {
|
||||
t.Fatalf("capability check does not contain representative rendered rules:\n%s", command)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildApplyCommandUsesAtomicReplacementBatch(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
batchPath := filepath.Join(dir, "batch.nft")
|
||||
nftPath := filepath.Join(dir, "nft")
|
||||
fake := `#!/bin/sh
|
||||
if [ "$1" = "list" ]; then
|
||||
exit 0
|
||||
fi
|
||||
if [ "$1" = "-f" ]; then
|
||||
cp "$2" "` + batchPath + `"
|
||||
fi
|
||||
exit 0
|
||||
`
|
||||
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
|
||||
t.Fatalf("write fake nft: %v", err)
|
||||
}
|
||||
|
||||
script := "table inet flvx {\n chain forward { }\n}"
|
||||
result := exec.Command("sh", "-c", buildApplyCommand(nftPath, script))
|
||||
if output, err := result.CombinedOutput(); err != nil {
|
||||
t.Fatalf("apply command failed: %v: %s", err, output)
|
||||
}
|
||||
batch, err := os.ReadFile(batchPath)
|
||||
if err != nil {
|
||||
t.Fatalf("read applied batch: %v", err)
|
||||
}
|
||||
want := "delete table inet flvx\n" + script + "\n"
|
||||
if string(batch) != want {
|
||||
t.Fatalf("atomic batch = %q, want %q", batch, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
|
||||
privateKey := mustGeneratePrivateKey(t)
|
||||
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
|
||||
|
||||
@@ -7,6 +7,10 @@ const (
|
||||
StatusPending = "pending"
|
||||
StatusApplied = "applied"
|
||||
StatusError = "error"
|
||||
|
||||
CounterDirectionDNAT = "dnat"
|
||||
CounterDirectionToTarget = "to-target"
|
||||
CounterDirectionFromTarget = "from-target"
|
||||
)
|
||||
|
||||
type Target struct {
|
||||
|
||||
@@ -16,6 +16,7 @@ type User struct {
|
||||
RoleID int `gorm:"column:role_id;not null"`
|
||||
ExpTime int64 `gorm:"column:exp_time;not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
FlowMiB int64 `gorm:"column:flow_mib;not null;default:0"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
@@ -31,24 +32,26 @@ func (User) TableName() string { return "user" }
|
||||
|
||||
// Forward maps to the "forward" table.
|
||||
type Forward struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
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"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
IPMaxConn int `gorm:"column:ip_max_conn;not null;default:0"`
|
||||
IPSpeedID sql.NullInt64 `gorm:"column:ip_speed_id"`
|
||||
ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"`
|
||||
ProxyProtocolReceive int `gorm:"column:proxy_protocol_receive;not null;default:0"`
|
||||
ProxyProtocolSend int `gorm:"column:proxy_protocol_send;not null;default:0"`
|
||||
}
|
||||
|
||||
func (Forward) TableName() string { return "forward" }
|
||||
@@ -131,6 +134,22 @@ type NftRuleBinding struct {
|
||||
|
||||
func (NftRuleBinding) TableName() string { return "nft_rule_binding" }
|
||||
|
||||
type NftCounterState struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||
Protocol string `gorm:"type:varchar(10);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||
Direction string `gorm:"type:varchar(20);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
|
||||
Bytes int64 `gorm:"not null;default:0"`
|
||||
Packets int64 `gorm:"not null;default:0"`
|
||||
CollectedTime int64 `gorm:"column:collected_time;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (NftCounterState) TableName() string { return "nft_counter_state" }
|
||||
|
||||
type SpeedLimit struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
@@ -212,6 +231,7 @@ type UserTunnel struct {
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
Num int `gorm:"not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
FlowMiB int64 `gorm:"column:flow_mib;not null;default:0"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
@@ -399,6 +419,7 @@ type UserBackup struct {
|
||||
RoleID int `json:"roleId"`
|
||||
ExpTime int64 `json:"expTime"`
|
||||
Flow int64 `json:"flow"`
|
||||
FlowMiB int64 `json:"flowMiB,omitempty"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
@@ -471,24 +492,26 @@ type ChainTunnelBackup struct {
|
||||
}
|
||||
|
||||
type ForwardBackup struct {
|
||||
ID int64 `json:"id"`
|
||||
UserID int64 `json:"userId"`
|
||||
UserName string `json:"userName"`
|
||||
Name string `json:"name"`
|
||||
TunnelID int64 `json:"tunnelId"`
|
||||
RemoteAddr string `json:"remoteAddr"`
|
||||
Strategy string `json:"strategy"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
Inx int `json:"inx"`
|
||||
SpeedID *int64 `json:"speedId,omitempty"`
|
||||
IPMaxConn int `json:"ipMaxConn,omitempty"`
|
||||
IPSpeedID *int64 `json:"ipSpeedId,omitempty"`
|
||||
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||
ProxyProtocol int `json:"proxyProtocol"`
|
||||
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"`
|
||||
ProxyProtocolReceive int `json:"proxyProtocolReceive,omitempty"`
|
||||
ProxyProtocolSend int `json:"proxyProtocolSend,omitempty"`
|
||||
}
|
||||
|
||||
type ForwardPortBackup struct {
|
||||
@@ -503,6 +526,7 @@ type UserTunnelBackup struct {
|
||||
SpeedID int64 `json:"speedId,omitempty"`
|
||||
Num int `json:"num"`
|
||||
Flow int64 `json:"flow"`
|
||||
FlowMiB int64 `json:"flowMiB,omitempty"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
@@ -577,19 +601,21 @@ type ImportResult struct {
|
||||
|
||||
// ForwardRecord is a minimal forward view used by control plane and flow policy.
|
||||
type ForwardRecord struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
ProxyProtocol int
|
||||
ID int64
|
||||
UserID int64
|
||||
UserName string
|
||||
Name string
|
||||
TunnelID int64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
ProxyProtocol int
|
||||
ProxyProtocolReceive int
|
||||
ProxyProtocolSend int
|
||||
}
|
||||
|
||||
// TunnelRecord is a minimal tunnel view used by control plane.
|
||||
@@ -684,6 +710,7 @@ type UserTunnelDetail struct {
|
||||
Status int
|
||||
TunnelFlow int
|
||||
Flow int64
|
||||
FlowMiB int64 `gorm:"column:flow_mib"`
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
Num int
|
||||
|
||||
@@ -14,15 +14,30 @@ var publicConfigKeys = map[string]struct{}{
|
||||
"app_logo": {},
|
||||
"app_favicon": {},
|
||||
"app_bg_image": {},
|
||||
"app_bg_image_light": {},
|
||||
"app_bg_image_dark": {},
|
||||
"cloudflare_site_key": {},
|
||||
"is_commercial": {},
|
||||
"hide_footer_brand": {},
|
||||
}
|
||||
|
||||
var sensitiveConfigKeys = map[string]struct{}{
|
||||
"jwt_secret": {},
|
||||
"license_key": {},
|
||||
"license_expiry": {},
|
||||
"license_machine_id": {},
|
||||
"machine_fingerprint": {},
|
||||
"cloudflare_secret_key": {},
|
||||
}
|
||||
|
||||
var systemManagedConfigKeys = map[string]struct{}{
|
||||
"license_key": {},
|
||||
"license_expiry": {},
|
||||
"license_machine_id": {},
|
||||
"is_commercial": {},
|
||||
"machine_fingerprint": {},
|
||||
}
|
||||
|
||||
func PolicyForConfig(name string) ConfigAccessPolicy {
|
||||
if IsPublicConfigKey(name) {
|
||||
return ConfigAccessPublic
|
||||
@@ -43,6 +58,11 @@ func IsSensitiveConfigKey(name string) bool {
|
||||
return ok
|
||||
}
|
||||
|
||||
func IsSystemManagedConfigKey(name string) bool {
|
||||
_, ok := systemManagedConfigKeys[normalizeConfigKey(name)]
|
||||
return ok
|
||||
}
|
||||
|
||||
func FilterSensitiveConfigs(in map[string]string) map[string]string {
|
||||
if len(in) == 0 {
|
||||
return map[string]string{}
|
||||
@@ -57,6 +77,20 @@ func FilterSensitiveConfigs(in map[string]string) map[string]string {
|
||||
return out
|
||||
}
|
||||
|
||||
func FilterBackupConfigs(in map[string]string) map[string]string {
|
||||
if len(in) == 0 {
|
||||
return map[string]string{}
|
||||
}
|
||||
out := make(map[string]string, len(in))
|
||||
for name, value := range in {
|
||||
if IsSensitiveConfigKey(name) || IsSystemManagedConfigKey(name) {
|
||||
continue
|
||||
}
|
||||
out[name] = value
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func normalizeConfigKey(name string) string {
|
||||
return strings.ToLower(strings.TrimSpace(name))
|
||||
}
|
||||
|
||||
@@ -12,7 +12,11 @@ func TestConfigPolicy(t *testing.T) {
|
||||
{name: "app_logo is public", key: "app_logo", want: ConfigAccessPublic},
|
||||
{name: "app_favicon is public", key: "app_favicon", want: ConfigAccessPublic},
|
||||
{name: "app_bg_image is public", key: "app_bg_image", want: ConfigAccessPublic},
|
||||
{name: "app_bg_image_light is public", key: "app_bg_image_light", want: ConfigAccessPublic},
|
||||
{name: "app_bg_image_dark is public", key: "app_bg_image_dark", want: ConfigAccessPublic},
|
||||
{name: "cloudflare_site_key is public", key: "cloudflare_site_key", want: ConfigAccessPublic},
|
||||
{name: "is_commercial is public", key: "is_commercial", want: ConfigAccessPublic},
|
||||
{name: "hide_footer_brand is public", key: "hide_footer_brand", want: ConfigAccessPublic},
|
||||
{name: "jwt_secret is sensitive", key: "jwt_secret", want: ConfigAccessSensitive},
|
||||
{name: "license_key is sensitive", key: "license_key", want: ConfigAccessSensitive},
|
||||
{name: "cloudflare_secret_key is sensitive", key: "cloudflare_secret_key", want: ConfigAccessSensitive},
|
||||
@@ -29,14 +33,14 @@ func TestConfigPolicy(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestConfigPolicyHelpers(t *testing.T) {
|
||||
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "cloudflare_site_key"}
|
||||
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "app_bg_image_light", "app_bg_image_dark", "cloudflare_site_key", "is_commercial", "hide_footer_brand"}
|
||||
for _, key := range publicKeys {
|
||||
if !IsPublicConfigKey(key) {
|
||||
t.Fatalf("expected %q to be public", key)
|
||||
}
|
||||
}
|
||||
|
||||
sensitiveKeys := []string{"jwt_secret", "license_key", "cloudflare_secret_key"}
|
||||
sensitiveKeys := []string{"jwt_secret", "license_key", "license_expiry", "license_machine_id", "machine_fingerprint", "cloudflare_secret_key"}
|
||||
for _, key := range sensitiveKeys {
|
||||
if !IsSensitiveConfigKey(key) {
|
||||
t.Fatalf("expected %q to be sensitive", key)
|
||||
@@ -67,3 +71,30 @@ func TestConfigPolicyHelpers(t *testing.T) {
|
||||
t.Fatal("expected cloudflare_secret_key to be filtered out")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFilterBackupConfigsOmitsLicenseState(t *testing.T) {
|
||||
filtered := FilterBackupConfigs(map[string]string{
|
||||
"app_name": "Brand",
|
||||
"license_key": "secret-license",
|
||||
"license_expiry": "never",
|
||||
"license_machine_id": "machine-id",
|
||||
"is_commercial": "true",
|
||||
"machine_fingerprint": "fingerprint",
|
||||
})
|
||||
if len(filtered) != 1 || filtered["app_name"] != "Brand" {
|
||||
t.Fatalf("unexpected backup configs: %+v", filtered)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemManagedConfigKeys(t *testing.T) {
|
||||
for _, key := range []string{"license_key", "license_expiry", "license_machine_id", "is_commercial", "machine_fingerprint"} {
|
||||
if !IsSystemManagedConfigKey(key) {
|
||||
t.Fatalf("expected %s to be system managed", key)
|
||||
}
|
||||
}
|
||||
for _, key := range []string{"jwt_secret", "cloudflare_secret_key", "app_name"} {
|
||||
if IsSystemManagedConfigKey(key) {
|
||||
t.Fatalf("did not expect %s to be system managed", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -104,6 +104,19 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
|
||||
return nil
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
return applyFlowUploadDeltasTx(tx, deltas)
|
||||
})
|
||||
}
|
||||
|
||||
func applyFlowUploadDeltasTx(tx *gorm.DB, deltas []FlowUploadCounterDelta) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
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))
|
||||
@@ -128,36 +141,34 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
|
||||
}
|
||||
}
|
||||
|
||||
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 _, 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 _, 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
|
||||
}
|
||||
}
|
||||
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
|
||||
})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ─── Open / Close ────────────────────────────────────────────────────
|
||||
@@ -282,6 +293,7 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
&model.Node{},
|
||||
&model.NodeSSHConfig{},
|
||||
&model.NftRuleBinding{},
|
||||
&model.NftCounterState{},
|
||||
&model.SpeedLimit{},
|
||||
&model.StatisticsFlow{},
|
||||
&model.Tunnel{},
|
||||
@@ -412,7 +424,7 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
}
|
||||
|
||||
if m.HasTable(&model.Forward{}) {
|
||||
for _, field := range []string{"MaxConn", "IPMaxConn", "IPSpeedID", "ProxyProtocol"} {
|
||||
for _, field := range []string{"MaxConn", "IPMaxConn", "IPSpeedID", "ProxyProtocol", "ProxyProtocolReceive", "ProxyProtocolSend"} {
|
||||
if m.HasColumn(&model.Forward{}, field) {
|
||||
continue
|
||||
}
|
||||
@@ -461,6 +473,14 @@ func seedData(db *gorm.DB) {
|
||||
|
||||
appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000}
|
||||
db.Where("id = ?", 1).FirstOrCreate(&appNameConfig)
|
||||
now := time.Now().UnixMilli()
|
||||
for name, value := range map[string]string{
|
||||
"is_commercial": "false",
|
||||
"hide_footer_brand": "false",
|
||||
} {
|
||||
cfg := model.ViteConfig{Name: name, Value: value, Time: now}
|
||||
db.Where("name = ?", name).FirstOrCreate(&cfg)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── User Queries ────────────────────────────────────────────────────
|
||||
@@ -588,6 +608,23 @@ func (r *Repository) UpsertConfig(name, value string, now int64) error {
|
||||
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertConfigs(values map[string]string, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
for name, value := range values {
|
||||
if err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "name"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{"value", "time"}),
|
||||
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// ─── Announcement Queries ────────────────────────────────────────────
|
||||
|
||||
func (r *Repository) GetAnnouncement() (*model.Announcement, error) {
|
||||
@@ -632,7 +669,7 @@ func (r *Repository) GetUserPackageTunnels(userID int64) ([]model.UserTunnelDeta
|
||||
}
|
||||
var items []model.UserTunnelDetail
|
||||
err := r.db.Model(&model.UserTunnel{}).
|
||||
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
|
||||
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.flow_mib, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
|
||||
Joins("LEFT JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id").
|
||||
Joins("LEFT JOIN speed_limit ON speed_limit.id = user_tunnel.speed_id").
|
||||
Where("user_tunnel.user_id = ?", userID).
|
||||
@@ -887,7 +924,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
||||
item := map[string]interface{}{
|
||||
"id": u.ID, "user": u.User, "name": u.User,
|
||||
"roleId": u.RoleID, "status": u.Status,
|
||||
"flow": u.Flow, "num": u.Num, "expTime": u.ExpTime,
|
||||
"flow": u.Flow, "flowMiB": u.FlowMiB, "num": u.Num, "expTime": u.ExpTime,
|
||||
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
|
||||
"updatedTime": nullableInt64(u.UpdatedTime),
|
||||
"inFlow": u.InFlow, "outFlow": u.OutFlow,
|
||||
@@ -932,31 +969,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
|
||||
IPMaxConn int
|
||||
IPSpeedID sql.NullInt64
|
||||
IPSpeedLimitName string
|
||||
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
|
||||
ProxyProtocolReceive int
|
||||
ProxyProtocolSend 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.ip_max_conn, forward.ip_speed_id, COALESCE(ip_speed_limit.name, '') AS ip_speed_limit_name, 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, forward.proxy_protocol_receive, forward.proxy_protocol_send").
|
||||
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").
|
||||
@@ -967,6 +1006,7 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
|
||||
items := make([]map[string]interface{}, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(row.ProxyProtocol, row.ProxyProtocolReceive, row.ProxyProtocolSend)
|
||||
inIP, inPort, err := resolveForwardIngress(r.db, row.ID, row.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -979,9 +1019,11 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
|
||||
"inFlow": row.InFlow, "outFlow": row.OutFlow,
|
||||
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
|
||||
"maxConn": row.MaxConn,
|
||||
"ipMaxConn": row.IPMaxConn,
|
||||
"proxyProtocol": row.ProxyProtocol,
|
||||
"maxConn": row.MaxConn,
|
||||
"ipMaxConn": row.IPMaxConn,
|
||||
"proxyProtocol": row.ProxyProtocol,
|
||||
"proxyProtocolReceive": proxyProtocolReceive,
|
||||
"proxyProtocolSend": proxyProtocolSend,
|
||||
}
|
||||
if row.SpeedID.Valid {
|
||||
item["speedId"] = row.SpeedID.Int64
|
||||
@@ -1950,7 +1992,7 @@ func (r *Repository) ExportAll() (*model.BackupData, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("export configs failed: %w", err)
|
||||
}
|
||||
backup.Configs = FilterSensitiveConfigs(configs)
|
||||
backup.Configs = FilterBackupConfigs(configs)
|
||||
|
||||
return backup, nil
|
||||
}
|
||||
@@ -2030,7 +2072,7 @@ func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("export configs failed: %w", err)
|
||||
}
|
||||
backup.Configs = FilterSensitiveConfigs(v)
|
||||
backup.Configs = FilterBackupConfigs(v)
|
||||
}
|
||||
return backup, nil
|
||||
}
|
||||
@@ -2052,7 +2094,7 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
|
||||
for _, u := range users {
|
||||
b := model.UserBackup{
|
||||
ID: u.ID, User: u.User, Pwd: u.Pwd, RoleID: u.RoleID,
|
||||
ExpTime: u.ExpTime, Flow: u.Flow, InFlow: u.InFlow, OutFlow: u.OutFlow,
|
||||
ExpTime: u.ExpTime, Flow: u.Flow, FlowMiB: u.FlowMiB, InFlow: u.InFlow, OutFlow: u.OutFlow,
|
||||
FlowResetTime: u.FlowResetTime, Num: u.Num,
|
||||
CreatedTime: u.CreatedTime, Status: u.Status,
|
||||
}
|
||||
@@ -2183,8 +2225,10 @@ 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,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
ProxyProtocolReceive: f.ProxyProtocolReceive,
|
||||
ProxyProtocolSend: f.ProxyProtocolSend,
|
||||
}
|
||||
if f.SpeedID.Valid {
|
||||
v := f.SpeedID.Int64
|
||||
@@ -2226,7 +2270,7 @@ func (r *Repository) exportUserTunnels() ([]model.UserTunnelBackup, error) {
|
||||
for _, ut := range uts {
|
||||
b := model.UserTunnelBackup{
|
||||
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
|
||||
Num: ut.Num, Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
Num: ut.Num, Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
FlowResetTime: ut.FlowResetTime, ExpTime: ut.ExpTime, Status: ut.Status,
|
||||
}
|
||||
if ut.SpeedID.Valid {
|
||||
@@ -2423,6 +2467,7 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
|
||||
RoleID: u.RoleID,
|
||||
ExpTime: u.ExpTime,
|
||||
Flow: u.Flow,
|
||||
FlowMiB: u.FlowMiB,
|
||||
InFlow: u.InFlow,
|
||||
OutFlow: u.OutFlow,
|
||||
FlowResetTime: u.FlowResetTime,
|
||||
@@ -2435,7 +2480,7 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
|
||||
err = tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user", "pwd", "role_id", "exp_time", "flow", "in_flow", "out_flow",
|
||||
"user", "pwd", "role_id", "exp_time", "flow", "flow_mib", "in_flow", "out_flow",
|
||||
"flow_reset_time", "num", "updated_time", "status", "password_changed_at",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
@@ -2618,29 +2663,31 @@ func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int
|
||||
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,
|
||||
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,
|
||||
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,
|
||||
ProxyProtocolReceive: f.ProxyProtocolReceive,
|
||||
ProxyProtocolSend: f.ProxyProtocolSend,
|
||||
}
|
||||
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", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol",
|
||||
"in_flow", "out_flow", "updated_time", "status", "inx", "speed_id", "ip_max_conn", "ip_speed_id", "proxy_protocol", "proxy_protocol_receive", "proxy_protocol_send",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
if err != nil {
|
||||
@@ -2671,6 +2718,7 @@ func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int6
|
||||
SpeedID: sql.NullInt64{Int64: ut.SpeedID, Valid: ut.SpeedID > 0},
|
||||
Num: ut.Num,
|
||||
Flow: ut.Flow,
|
||||
FlowMiB: ut.FlowMiB,
|
||||
InFlow: ut.InFlow,
|
||||
OutFlow: ut.OutFlow,
|
||||
FlowResetTime: ut.FlowResetTime,
|
||||
@@ -2680,7 +2728,7 @@ func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int6
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user_id", "tunnel_id", "speed_id", "num", "flow", "in_flow", "out_flow",
|
||||
"user_id", "tunnel_id", "speed_id", "num", "flow", "flow_mib", "in_flow", "out_flow",
|
||||
"flow_reset_time", "exp_time", "status",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
@@ -2818,7 +2866,7 @@ func importPermissions(tx *gorm.DB, permissions []model.PermissionBackup, _ int6
|
||||
}
|
||||
|
||||
func importConfigs(tx *gorm.DB, configs map[string]string, now int64) (int, error) {
|
||||
configs = FilterSensitiveConfigs(configs)
|
||||
configs = FilterBackupConfigs(configs)
|
||||
count := 0
|
||||
for name, value := range configs {
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
|
||||
@@ -97,6 +97,10 @@ func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
|
||||
seedConfig(t, r, "cloudflare_site_key", "site-key")
|
||||
seedConfig(t, r, "jwt_secret", "jwt-secret")
|
||||
seedConfig(t, r, "license_key", "license-secret")
|
||||
seedConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
|
||||
seedConfig(t, r, "license_machine_id", "machine-id")
|
||||
seedConfig(t, r, "is_commercial", "true")
|
||||
seedConfig(t, r, "machine_fingerprint", "machine-fingerprint")
|
||||
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-secret")
|
||||
|
||||
for _, tc := range []struct {
|
||||
@@ -117,7 +121,7 @@ func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
|
||||
if backup.Configs["cloudflare_site_key"] != "site-key" {
|
||||
t.Fatalf("expected public config in export, got %+v", backup.Configs)
|
||||
}
|
||||
for _, key := range []string{"jwt_secret", "license_key", "cloudflare_secret_key"} {
|
||||
for _, key := range []string{"jwt_secret", "license_key", "license_expiry", "license_machine_id", "is_commercial", "machine_fingerprint", "cloudflare_secret_key"} {
|
||||
if _, ok := backup.Configs[key]; ok {
|
||||
t.Fatalf("expected %s to be omitted from export, got %+v", key, backup.Configs)
|
||||
}
|
||||
@@ -136,12 +140,20 @@ func TestImportIgnoresSensitiveConfigs(t *testing.T) {
|
||||
seedConfig(t, r, "app_name", "before")
|
||||
seedConfig(t, r, "jwt_secret", "jwt-before")
|
||||
seedConfig(t, r, "license_key", "license-before")
|
||||
seedConfig(t, r, "license_expiry", "expiry-before")
|
||||
seedConfig(t, r, "license_machine_id", "machine-before")
|
||||
seedConfig(t, r, "is_commercial", "true")
|
||||
seedConfig(t, r, "machine_fingerprint", "fingerprint-before")
|
||||
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
|
||||
backup := &model.BackupData{Configs: map[string]string{
|
||||
"app_name": "after",
|
||||
"jwt_secret": "jwt-after",
|
||||
"license_key": "license-after",
|
||||
"license_expiry": "expiry-after",
|
||||
"license_machine_id": "machine-after",
|
||||
"is_commercial": "false",
|
||||
"machine_fingerprint": "fingerprint-after",
|
||||
"cloudflare_secret_key": "cloudflare-after",
|
||||
}}
|
||||
|
||||
@@ -156,6 +168,10 @@ func TestImportIgnoresSensitiveConfigs(t *testing.T) {
|
||||
assertConfigValue(t, r, "app_name", "after")
|
||||
assertConfigValue(t, r, "jwt_secret", "jwt-before")
|
||||
assertConfigValue(t, r, "license_key", "license-before")
|
||||
assertConfigValue(t, r, "license_expiry", "expiry-before")
|
||||
assertConfigValue(t, r, "license_machine_id", "machine-before")
|
||||
assertConfigValue(t, r, "is_commercial", "true")
|
||||
assertConfigValue(t, r, "machine_fingerprint", "fingerprint-before")
|
||||
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
|
||||
}
|
||||
|
||||
|
||||
@@ -44,20 +44,23 @@ func (r *Repository) ListForwardsByTunnelTx(tx *gorm.DB, tunnelID int64) ([]mode
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
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,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
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,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
|
||||
@@ -11,6 +11,8 @@ import (
|
||||
|
||||
type FlowUploadForwardMeta struct {
|
||||
ForwardID int64
|
||||
UserID int64
|
||||
UserTunnelID int64
|
||||
TunnelID int64
|
||||
TrafficRatio float64
|
||||
TunnelFlow int64
|
||||
@@ -59,6 +61,8 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
|
||||
type row struct {
|
||||
ForwardID int64 `gorm:"column:forward_id"`
|
||||
UserID int64 `gorm:"column:user_id"`
|
||||
UserTunnelID int64 `gorm:"column:user_tunnel_id"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id"`
|
||||
TrafficRatio float64 `gorm:"column:traffic_ratio"`
|
||||
TunnelFlow int64 `gorm:"column:tunnel_flow"`
|
||||
@@ -68,8 +72,9 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
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").
|
||||
Select("f.id AS forward_id, f.user_id AS user_id, COALESCE(ut.id, 0) AS user_tunnel_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").
|
||||
Joins("LEFT JOIN user_tunnel ut ON ut.user_id = f.user_id AND ut.tunnel_id = f.tunnel_id").
|
||||
Where("f.id IN ?", chunk).
|
||||
Scan(&rows).Error
|
||||
if err != nil {
|
||||
@@ -84,6 +89,8 @@ func (r *Repository) GetFlowUploadForwardMetas(forwardIDs []int64) (map[int64]Fl
|
||||
}
|
||||
out[row.ForwardID] = FlowUploadForwardMeta{
|
||||
ForwardID: row.ForwardID,
|
||||
UserID: row.UserID,
|
||||
UserTunnelID: row.UserTunnelID,
|
||||
TunnelID: row.TunnelID,
|
||||
TrafficRatio: row.TrafficRatio,
|
||||
TunnelFlow: row.TunnelFlow,
|
||||
@@ -113,20 +120,23 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
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,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
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,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -148,20 +158,23 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
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,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
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,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -183,20 +196,23 @@ func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]mode
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
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,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
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,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -219,20 +235,23 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
fr := model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
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,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
}
|
||||
if strings.TrimSpace(fr.Strategy) == "" {
|
||||
fr.Strategy = "fifo"
|
||||
|
||||
@@ -64,7 +64,7 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("get metas: %v", err)
|
||||
}
|
||||
if metas[20].TunnelID != 1 || metas[20].TrafficRatio != 2 || metas[20].TunnelFlow != 3 {
|
||||
if metas[20].UserID != 2 || metas[20].UserTunnelID != 10 || 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 {
|
||||
@@ -86,6 +86,99 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNftTrafficAccountingAppliesFlowQuotaAndStates(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
seedFlowBatchRows(t, r, nowMs)
|
||||
|
||||
quotaViews, err := r.ApplyNftTrafficAccounting(
|
||||
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
|
||||
map[int64]int64{2: 1140},
|
||||
[]NftCounterStateInput{{
|
||||
NodeID: 11,
|
||||
ForwardID: 20,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1400,
|
||||
Packets: 14,
|
||||
CollectedTime: nowMs,
|
||||
}},
|
||||
now,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("ApplyNftTrafficAccounting: %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)
|
||||
}
|
||||
if quotaViews[2] == nil || quotaViews[2].DailyUsedBytes != 1140 || quotaViews[2].MonthlyUsedBytes != 1140 {
|
||||
t.Fatalf("unexpected quota view: %#v", quotaViews[2])
|
||||
}
|
||||
states, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(states) != 1 || states[0].ForwardID != 20 || states[0].Bytes != 1400 {
|
||||
t.Fatalf("unexpected nft counter state: %+v", states)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyNftTrafficAccountingRollsBackFlowAndQuotaWhenStateWriteFails(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting-rollback.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now()
|
||||
nowMs := now.UnixMilli()
|
||||
seedFlowBatchRows(t, r, nowMs)
|
||||
if err := r.DB().Exec(`DROP TABLE nft_counter_state`).Error; err != nil {
|
||||
t.Fatalf("drop nft_counter_state: %v", err)
|
||||
}
|
||||
|
||||
_, err = r.ApplyNftTrafficAccounting(
|
||||
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
|
||||
map[int64]int64{2: 1140},
|
||||
[]NftCounterStateInput{{
|
||||
NodeID: 11,
|
||||
ForwardID: 20,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 1400,
|
||||
Packets: 14,
|
||||
CollectedTime: nowMs,
|
||||
}},
|
||||
now,
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatalf("expected ApplyNftTrafficAccounting to fail")
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 0 {
|
||||
t.Fatalf("expected forward flow rollback, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 0 {
|
||||
t.Fatalf("expected user flow rollback, got %d", got)
|
||||
}
|
||||
if got := mustFlowBatchCount(t, r, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 2), 0)`); got != 0 {
|
||||
t.Fatalf("expected quota rollback, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db"))
|
||||
if err != nil {
|
||||
@@ -106,7 +199,7 @@ func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T
|
||||
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 {
|
||||
if meta.ForwardID != 25 || meta.UserID != 2 || meta.UserTunnelID != 0 || meta.TunnelID != 99 || meta.TrafficRatio != 1 || meta.TunnelFlow != 1 {
|
||||
t.Fatalf("unexpected fallback meta: %#v", meta)
|
||||
}
|
||||
}
|
||||
@@ -167,3 +260,19 @@ func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...inter
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func seedFlowBatchRows(t *testing.T, r *Repository, now int64) {
|
||||
t.Helper()
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResetForwardFlowOnlyUpdatesSelectedForward(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "forward-flow-reset.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
const originalUpdated int64 = 1000
|
||||
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, 'owner', 'pwd', 1, 0, 100, 700, 900, 0, 10, 1000, 1000, 1)
|
||||
`).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, 'tunnel', 1, 1, 'tls', 1, 1000, 1000, 1, NULL, 0)
|
||||
`).Error; err != nil {
|
||||
t.Fatalf("insert 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(10, 2, 1, 10, 100, 500, 600, 0, 0, 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, 'owner', 'target', 1, '127.0.0.1:80', 'fifo', 111, 222, 1000, ?, 1, 0),
|
||||
(21, 2, 'owner', 'other', 1, '127.0.0.1:81', 'fifo', 333, 444, 1000, ?, 1, 1)
|
||||
`, originalUpdated, originalUpdated).Error; err != nil {
|
||||
t.Fatalf("insert forwards: %v", err)
|
||||
}
|
||||
|
||||
const resetAt int64 = 2000
|
||||
if err := r.ResetForwardFlow(20, resetAt); err != nil {
|
||||
t.Fatalf("ResetForwardFlow: %v", err)
|
||||
}
|
||||
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM forward WHERE id = 20", 0)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM forward WHERE id = 20", 0)
|
||||
assertForwardFlowResetValue(t, r, "SELECT updated_time FROM forward WHERE id = 20", resetAt)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM forward WHERE id = 21", 333)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM forward WHERE id = 21", 444)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM user WHERE id = 2", 700)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM user WHERE id = 2", 900)
|
||||
assertForwardFlowResetValue(t, r, "SELECT in_flow FROM user_tunnel WHERE id = 10", 500)
|
||||
assertForwardFlowResetValue(t, r, "SELECT out_flow FROM user_tunnel WHERE id = 10", 600)
|
||||
}
|
||||
|
||||
func TestResetForwardFlowRejectsUninitializedRepository(t *testing.T) {
|
||||
var r *Repository
|
||||
if err := r.ResetForwardFlow(20, 2000); err == nil {
|
||||
t.Fatal("expected uninitialized repository error")
|
||||
}
|
||||
}
|
||||
|
||||
func assertForwardFlowResetValue(t *testing.T, r *Repository, query string, want int64) {
|
||||
t.Helper()
|
||||
var got int64
|
||||
if err := r.DB().Raw(query).Scan(&got).Error; err != nil {
|
||||
t.Fatalf("query %q: %v", query, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("query %q returned %d, want %d", query, got, want)
|
||||
}
|
||||
}
|
||||
@@ -160,7 +160,7 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
|
||||
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)
|
||||
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, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("CreateForwardTx: %v", err)
|
||||
}
|
||||
@@ -175,7 +175,7 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
|
||||
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 {
|
||||
if err := r.UpdateForward(forwardID, "per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 9, int64(22), 0, 0, 0); err != nil {
|
||||
t.Fatalf("UpdateForward: %v", err)
|
||||
}
|
||||
record, err = r.GetForwardRecord(forwardID)
|
||||
@@ -216,6 +216,83 @@ func TestForwardRepositoryPersistsPerIPLimits(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardRepositoryPersistsProxyProtocolReceiveAndSend(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", "proxy-protocol-forward", 2, "1.1.1.1:443", "fifo", now, 1, []int64{3}, 24000, "", nil, 0, 0, nil, 0, 1, 2)
|
||||
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.ProxyProtocolReceive != 1 || record.ProxyProtocolSend != 2 {
|
||||
t.Fatalf("expected proxyProtocol receive/send 1/2 after create, got %d/%d", record.ProxyProtocolReceive, record.ProxyProtocolSend)
|
||||
}
|
||||
|
||||
if err := r.UpdateForward(forwardID, "proxy-protocol-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0, 2, 1); err != nil {
|
||||
t.Fatalf("UpdateForward: %v", err)
|
||||
}
|
||||
record, err = r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetForwardRecord after update: %v", err)
|
||||
}
|
||||
if record.ProxyProtocolReceive != 2 || record.ProxyProtocolSend != 1 {
|
||||
t.Fatalf("expected proxyProtocol receive/send 2/1 after update, got %d/%d", record.ProxyProtocolReceive, record.ProxyProtocolSend)
|
||||
}
|
||||
|
||||
records, err := r.ListForwardsByTunnel(2)
|
||||
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].ProxyProtocolReceive != 2 || records[0].ProxyProtocolSend != 1 {
|
||||
t.Fatalf("expected listed proxyProtocol receive/send 2/1, got %d/%d", records[0].ProxyProtocolReceive, records[0].ProxyProtocolSend)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardRepositoryMapsLegacyProxyProtocolToSend(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: "legacy-proxy-protocol-forward",
|
||||
TunnelID: 9,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
ProxyProtocol: 2,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create legacy forward: %v", err)
|
||||
}
|
||||
forwardID := mustRepoLastInsertID(t, r)
|
||||
|
||||
record, err := r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetForwardRecord: %v", err)
|
||||
}
|
||||
if record.ProxyProtocolReceive != 0 || record.ProxyProtocolSend != 2 {
|
||||
t.Fatalf("expected legacy proxyProtocol to map to receive/send 0/2, got %d/%d", record.ProxyProtocolReceive, record.ProxyProtocolSend)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackForwardFieldsRestoresPerIPLimits(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
@@ -224,15 +301,15 @@ func TestRollbackForwardFieldsRestoresPerIPLimits(t *testing.T) {
|
||||
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)
|
||||
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, 0, 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 {
|
||||
if err := r.UpdateForward(forwardID, "rollback-per-ip-forward", 2, "2.2.2.2:443", "fifo", now+1, nil, 0, 0, nil, 0, 0, 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)
|
||||
r.RollbackForwardFields(forwardID, 1, "admin", "rollback-per-ip-forward", 2, "1.1.1.1:443", "fifo", 1, nil, 7, 5, int64(21), 2, 0, 2, now+2)
|
||||
|
||||
record, err := r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
|
||||
@@ -132,3 +132,31 @@ func TestUpsertTunnelMetricBucketsIsSafeUnderConcurrency(t *testing.T) {
|
||||
t.Fatalf("expected bytesOut %d, got %d", wantOut, rows[0].BytesOut)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetLatestTunnelQualitiesIncludesChainDetails(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
if err := r.InsertTunnelQuality(&model.TunnelQuality{
|
||||
TunnelID: 7,
|
||||
Timestamp: time.Now().UnixMilli(),
|
||||
Success: 1,
|
||||
ChainDetails: `{"primaryPath":[],"candidateHops":[{"fromNodeId":10,"toNodeId":31}]}`,
|
||||
}); err != nil {
|
||||
t.Fatalf("insert tunnel quality: %v", err)
|
||||
}
|
||||
|
||||
items, err := r.GetLatestTunnelQualities()
|
||||
if err != nil {
|
||||
t.Fatalf("get latest tunnel qualities: %v", err)
|
||||
}
|
||||
if len(items) != 1 {
|
||||
t.Fatalf("expected one latest tunnel quality, got %+v", items)
|
||||
}
|
||||
if items[0].ChainDetails == "" {
|
||||
t.Fatalf("expected chain details in latest quality row, got %+v", items[0])
|
||||
}
|
||||
}
|
||||
|
||||
@@ -37,7 +37,14 @@ func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool
|
||||
return cnt > 0, err
|
||||
}
|
||||
|
||||
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64) (int64, error) {
|
||||
func optionalFlowMiB(values []int64) int64 {
|
||||
if len(values) > 0 {
|
||||
return values[0]
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64, flowMiB ...int64) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
@@ -47,6 +54,7 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
|
||||
RoleID: roleID,
|
||||
ExpTime: expTime,
|
||||
Flow: flow,
|
||||
FlowMiB: optionalFlowMiB(flowMiB),
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
@@ -75,7 +83,7 @@ func (r *Repository) GetUserRoleID(userID int64) (int, error) {
|
||||
return user.RoleID, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
|
||||
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -85,6 +93,7 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
|
||||
"user": username,
|
||||
"pwd": pwdHash,
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -95,7 +104,7 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
|
||||
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -104,6 +113,7 @@ func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow i
|
||||
Updates(map[string]interface{}{
|
||||
"user": username,
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -126,7 +136,7 @@ func (r *Repository) UpdateUserPassword(userID int64, pwdHash string, now int64)
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) {
|
||||
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64, flowMiB ...int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
@@ -134,6 +144,7 @@ func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num in
|
||||
Where("user_id = ?", userID).
|
||||
Updates(map[string]interface{}{
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -197,6 +208,19 @@ func (r *Repository) ResetUserFlowByUserTunnel(userTunnelID int64) {
|
||||
Updates(map[string]interface{}{"in_flow": 0, "out_flow": 0}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ResetForwardFlow(forwardID int64, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", forwardID).
|
||||
Updates(map[string]interface{}{
|
||||
"in_flow": 0,
|
||||
"out_flow": 0,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) GetUsernameByID(userID int64) string {
|
||||
if r == nil || r.db == nil {
|
||||
return ""
|
||||
@@ -637,7 +661,7 @@ func (r *Repository) DeleteUserTunnel(id int64) error {
|
||||
return r.db.Where("id = ?", id).Delete(&model.UserTunnel{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int) error {
|
||||
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -645,6 +669,7 @@ func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, fl
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -679,7 +704,7 @@ func (r *Repository) GetExistingUserTunnel(userID, tunnelID int64) (id int64, fl
|
||||
return ut.ID, ut.Flow, int64(ut.Num), ut.ExpTime, ut.FlowResetTime, ut.SpeedID, ut.Status, nil
|
||||
}
|
||||
|
||||
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int) error {
|
||||
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -689,6 +714,7 @@ func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
Num: num,
|
||||
Flow: flow,
|
||||
FlowMiB: optionalFlowMiB(flowMiB),
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
@@ -698,7 +724,7 @@ func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{
|
||||
return r.db.Create(&ut).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int) error {
|
||||
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -707,6 +733,7 @@ func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow
|
||||
Updates(map[string]interface{}{
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -726,23 +753,26 @@ 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, ipMaxConn int, ipSpeedID interface{}, 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, proxyProtocolReceive int, proxyProtocolSend int) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
_, proxyProtocolSend = normalizeForwardProxyProtocol(proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
return r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"updated_time": now,
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"proxy_protocol_receive": proxyProtocolReceive,
|
||||
"proxy_protocol_send": proxyProtocolSend,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -769,6 +799,9 @@ func (r *Repository) DeleteForwardCascade(forwardID int64) error {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.Where("forward_id = ?", forwardID).Delete(&model.ForwardPort{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -816,26 +849,29 @@ 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, ipMaxConn int, ipSpeedID interface{}, 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, proxyProtocolReceive int, proxyProtocolSend int, now int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
_, proxyProtocolSend = normalizeForwardProxyProtocol(proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
_ = 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,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"updated_time": now,
|
||||
"user_id": userID,
|
||||
"user_name": userName,
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"status": status,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"ip_max_conn": ipMaxConn,
|
||||
"ip_speed_id": nullInt64FromInterface(ipSpeedID),
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"proxy_protocol_receive": proxyProtocolReceive,
|
||||
"proxy_protocol_send": proxyProtocolSend,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -1271,7 +1307,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
return 0, false, err
|
||||
}
|
||||
var user model.User
|
||||
if err := r.db.Select("flow, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
if err := r.db.Select("flow, flow_mib, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
flow := user.Flow
|
||||
@@ -1283,6 +1319,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
TunnelID: tunnelID,
|
||||
Num: num,
|
||||
Flow: flow,
|
||||
FlowMiB: user.FlowMiB,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowReset,
|
||||
@@ -1295,30 +1332,33 @@ 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, ipMaxConn int, ipSpeedID interface{}, 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, proxyProtocolReceive int, proxyProtocolSend int) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
_, proxyProtocolSend = normalizeForwardProxyProtocol(proxyProtocol, proxyProtocolReceive, proxyProtocolSend)
|
||||
var forwardID int64
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
fwd := model.Forward{
|
||||
UserID: userID,
|
||||
UserName: userName,
|
||||
Name: name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: remoteAddr,
|
||||
Strategy: strategy,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
IPMaxConn: ipMaxConn,
|
||||
IPSpeedID: nullInt64FromInterface(ipSpeedID),
|
||||
ProxyProtocol: proxyProtocol,
|
||||
UserID: userID,
|
||||
UserName: userName,
|
||||
Name: name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: remoteAddr,
|
||||
Strategy: strategy,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
IPMaxConn: ipMaxConn,
|
||||
IPSpeedID: nullInt64FromInterface(ipSpeedID),
|
||||
ProxyProtocol: proxyProtocol,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
}
|
||||
if err := tx.Create(&fwd).Error; err != nil {
|
||||
return err
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const (
|
||||
nftCounterProtocolTCP = "tcp"
|
||||
nftCounterProtocolUDP = "udp"
|
||||
|
||||
nftCounterDirectionToTarget = "to-target"
|
||||
nftCounterDirectionFromTarget = "from-target"
|
||||
)
|
||||
|
||||
type NftCounterStateInput struct {
|
||||
NodeID int64
|
||||
ForwardID int64
|
||||
Protocol string
|
||||
Direction string
|
||||
RuleHash string
|
||||
Bytes uint64
|
||||
Packets uint64
|
||||
CollectedTime int64
|
||||
}
|
||||
|
||||
type NftablesCollectionNode struct {
|
||||
NodeID int64
|
||||
Config model.NodeSSHConfig
|
||||
}
|
||||
|
||||
func (r *Repository) ListNftablesNodesForCollection() ([]NftablesCollectionNode, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
type collectionRow struct {
|
||||
NodeID int64 `gorm:"column:node_id"`
|
||||
ConfigID int64 `gorm:"column:config_id"`
|
||||
Host string `gorm:"column:host"`
|
||||
Port int `gorm:"column:port"`
|
||||
Username string `gorm:"column:username"`
|
||||
AuthType string `gorm:"column:auth_type"`
|
||||
Password string `gorm:"column:password"`
|
||||
PrivateKey string `gorm:"column:private_key"`
|
||||
Passphrase string `gorm:"column:passphrase"`
|
||||
SudoMode string `gorm:"column:sudo_mode"`
|
||||
CreatedTime int64 `gorm:"column:created_time"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time"`
|
||||
}
|
||||
|
||||
var rows []collectionRow
|
||||
if err := r.db.Table("node").
|
||||
Select("node.id AS node_id, node_ssh_config.id AS config_id, node_ssh_config.host, node_ssh_config.port, node_ssh_config.username, node_ssh_config.auth_type, node_ssh_config.password, node_ssh_config.private_key, node_ssh_config.passphrase, node_ssh_config.sudo_mode, node_ssh_config.created_time, node_ssh_config.updated_time").
|
||||
Joins("JOIN node_ssh_config ON node_ssh_config.node_id = node.id").
|
||||
Where("node.status = ? AND LOWER(TRIM(node.forward_mode)) = ?", 1, "nftables").
|
||||
Order("node.id ASC").
|
||||
Scan(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
nodes := make([]NftablesCollectionNode, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
nodes = append(nodes, NftablesCollectionNode{
|
||||
NodeID: row.NodeID,
|
||||
Config: model.NodeSSHConfig{
|
||||
ID: row.ConfigID,
|
||||
NodeID: row.NodeID,
|
||||
Host: row.Host,
|
||||
Port: row.Port,
|
||||
Username: row.Username,
|
||||
AuthType: row.AuthType,
|
||||
Password: nullStringFromInterface(row.Password),
|
||||
PrivateKey: nullStringFromInterface(row.PrivateKey),
|
||||
Passphrase: nullStringFromInterface(row.Passphrase),
|
||||
SudoMode: row.SudoMode,
|
||||
CreatedTime: row.CreatedTime,
|
||||
UpdatedTime: row.UpdatedTime,
|
||||
},
|
||||
})
|
||||
}
|
||||
return nodes, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetNftCounterStatesByNode(nodeID int64) ([]model.NftCounterState, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var rows []model.NftCounterState
|
||||
err := r.db.Where("node_id = ?", nodeID).
|
||||
Order("forward_id ASC, protocol ASC, direction ASC").
|
||||
Find(&rows).Error
|
||||
return rows, err
|
||||
}
|
||||
|
||||
func (r *Repository) UpsertNftCounterStates(inputs []NftCounterStateInput, now int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if len(inputs) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return r.db.Transaction(func(tx *gorm.DB) error {
|
||||
return upsertNftCounterStatesTx(tx, inputs, now)
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) ApplyNftTrafficAccounting(deltas []FlowUploadCounterDelta, quotaUsage map[int64]int64, states []NftCounterStateInput, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
|
||||
quotaViews := map[int64]*model.UserQuotaView{}
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
if err := applyFlowUploadDeltasTx(tx, deltas); err != nil {
|
||||
return err
|
||||
}
|
||||
var err error
|
||||
quotaViews, err = r.addUserQuotaUsageBatchTx(tx, quotaUsage, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return upsertNftCounterStatesTx(tx, states, now.UnixMilli())
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return quotaViews, nil
|
||||
}
|
||||
|
||||
func upsertNftCounterStatesTx(tx *gorm.DB, inputs []NftCounterStateInput, now int64) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
for _, input := range inputs {
|
||||
row, ok := nftCounterStateFromInput(input, now)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{
|
||||
{Name: "node_id"},
|
||||
{Name: "forward_id"},
|
||||
{Name: "protocol"},
|
||||
{Name: "direction"},
|
||||
},
|
||||
DoUpdates: clause.Assignments(map[string]interface{}{
|
||||
"rule_hash": row.RuleHash,
|
||||
"bytes": row.Bytes,
|
||||
"packets": row.Packets,
|
||||
"collected_time": row.CollectedTime,
|
||||
"updated_time": row.UpdatedTime,
|
||||
}),
|
||||
}).Create(&row).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Repository) DeleteNftCounterStatesByForward(forwardID int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error
|
||||
}
|
||||
|
||||
func nftCounterStateFromInput(input NftCounterStateInput, now int64) (model.NftCounterState, bool) {
|
||||
protocol := strings.ToLower(strings.TrimSpace(input.Protocol))
|
||||
direction := strings.ToLower(strings.TrimSpace(input.Direction))
|
||||
if input.NodeID <= 0 || input.ForwardID <= 0 || !isValidNftCounterProtocol(protocol) || !isValidNftCounterDirection(direction) {
|
||||
return model.NftCounterState{}, false
|
||||
}
|
||||
if input.Bytes > uint64(math.MaxInt64) || input.Packets > uint64(math.MaxInt64) {
|
||||
return model.NftCounterState{}, false
|
||||
}
|
||||
return model.NftCounterState{
|
||||
NodeID: input.NodeID,
|
||||
ForwardID: input.ForwardID,
|
||||
Protocol: protocol,
|
||||
Direction: direction,
|
||||
RuleHash: strings.TrimSpace(input.RuleHash),
|
||||
Bytes: int64(input.Bytes),
|
||||
Packets: int64(input.Packets),
|
||||
CollectedTime: input.CollectedTime,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
}, true
|
||||
}
|
||||
|
||||
func isValidNftCounterProtocol(protocol string) bool {
|
||||
return protocol == nftCounterProtocolTCP || protocol == nftCounterProtocolUDP
|
||||
}
|
||||
|
||||
func isValidNftCounterDirection(direction string) bool {
|
||||
return direction == nftCounterDirectionToTarget || direction == nftCounterDirectionFromTarget
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"math"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
first := []NftCounterStateInput{
|
||||
{
|
||||
NodeID: 11,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-a",
|
||||
Bytes: 100,
|
||||
Packets: 10,
|
||||
CollectedTime: 1000,
|
||||
},
|
||||
{
|
||||
NodeID: 0,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
Bytes: 999,
|
||||
},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(first, 2000); err != nil {
|
||||
t.Fatalf("first UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
second := []NftCounterStateInput{
|
||||
{
|
||||
NodeID: 11,
|
||||
ForwardID: 42,
|
||||
Protocol: "tcp",
|
||||
Direction: "to-target",
|
||||
RuleHash: "hash-b",
|
||||
Bytes: 250,
|
||||
Packets: 25,
|
||||
CollectedTime: 3000,
|
||||
},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(second, 4000); err != nil {
|
||||
t.Fatalf("second UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected one counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
got := rows[0]
|
||||
if got.ForwardID != 42 || got.Protocol != "tcp" || got.Direction != "to-target" {
|
||||
t.Fatalf("unexpected counter state key: %+v", got)
|
||||
}
|
||||
if got.RuleHash != "hash-b" || got.Bytes != 250 || got.Packets != 25 || got.CollectedTime != 3000 {
|
||||
t.Fatalf("counter state was not updated: %+v", got)
|
||||
}
|
||||
if got.CreatedTime != 2000 || got.UpdatedTime != 4000 {
|
||||
t.Fatalf("unexpected timestamps after upsert: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateDeleteByForwardRemovesOnlyMatchingRows(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||
{NodeID: 12, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "c", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
if err := r.DeleteNftCounterStatesByForward(42); err != nil {
|
||||
t.Fatalf("DeleteNftCounterStatesByForward: %v", err)
|
||||
}
|
||||
|
||||
node11, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode(11): %v", err)
|
||||
}
|
||||
if len(node11) != 1 || node11[0].ForwardID != 43 {
|
||||
t.Fatalf("expected only forward 43 for node 11, got %+v", node11)
|
||||
}
|
||||
node12, err := r.GetNftCounterStatesByNode(12)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode(12): %v", err)
|
||||
}
|
||||
if len(node12) != 0 {
|
||||
t.Fatalf("expected forward 42 state removed from node 12, got %+v", node12)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteForwardCascadeRemovesNftCounterStateOnlyForDeletedForward(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwards := []model.Forward{
|
||||
{ID: 42, UserID: 1, UserName: "admin", Name: "forward-a", TunnelID: 10, RemoteAddr: "203.0.113.1:80", Strategy: "fifo", CreatedTime: now, UpdatedTime: now, Status: 1},
|
||||
{ID: 43, UserID: 1, UserName: "admin", Name: "forward-b", TunnelID: 10, RemoteAddr: "203.0.113.2:80", Strategy: "fifo", CreatedTime: now, UpdatedTime: now, Status: 1},
|
||||
}
|
||||
if err := r.DB().Create(&forwards).Error; err != nil {
|
||||
t.Fatalf("seed forwards: %v", err)
|
||||
}
|
||||
if err := r.UpsertNftCounterStates([]NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: now},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: now},
|
||||
}, now); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
if err := r.DeleteForwardCascade(42); err != nil {
|
||||
t.Fatalf("DeleteForwardCascade: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 || rows[0].ForwardID != 43 {
|
||||
t.Fatalf("expected only forward 43 counter state to remain, got %+v", rows)
|
||||
}
|
||||
var deletedForwardCount int64
|
||||
if err := r.DB().Model(&model.Forward{}).Where("id = ?", int64(42)).Count(&deletedForwardCount).Error; err != nil {
|
||||
t.Fatalf("count deleted forward: %v", err)
|
||||
}
|
||||
if deletedForwardCount != 0 {
|
||||
t.Fatalf("expected forward 42 deleted, count=%d", deletedForwardCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateUpsertSkipsInvalidProtocolAndDirection(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "icmp", Direction: "to-target", RuleHash: "bad-protocol", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "tcp", Direction: "sideways", RuleHash: "bad-direction", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 44, Protocol: " UDP ", Direction: " FROM-TARGET ", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
if rows[0].ForwardID != 44 || rows[0].Protocol != "udp" || rows[0].Direction != "from-target" {
|
||||
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftCounterStateUpsertSkipsCountersAboveInt64(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
inputs := []NftCounterStateInput{
|
||||
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "too-large", Bytes: uint64(math.MaxInt64) + 1, Packets: 10, CollectedTime: 1000},
|
||||
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||
}
|
||||
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||
}
|
||||
|
||||
rows, err := r.GetNftCounterStatesByNode(11)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||
}
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||
}
|
||||
if rows[0].ForwardID != 43 || rows[0].Bytes != 300 || rows[0].Packets != 30 {
|
||||
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestListNftablesNodesForCollectionReturnsActiveNftablesWithSSHOrdered(t *testing.T) {
|
||||
r, err := Open(filepath.Join(t.TempDir(), "nft-collection.db"))
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
seedCollectionNode(t, r, 1, "agent", 1, now)
|
||||
seedCollectionNode(t, r, 2, " nftables ", 1, now)
|
||||
seedCollectionNode(t, r, 3, "NFTABLES", 0, now)
|
||||
seedCollectionNode(t, r, 4, "nftables", 1, now)
|
||||
seedCollectionNode(t, r, 5, "nftables", 1, now)
|
||||
|
||||
if err := r.UpsertNodeSSHConfig(4, NftSSHConfigInput{
|
||||
Host: "203.0.113.4",
|
||||
Port: 2222,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "secret-4",
|
||||
SudoMode: "none",
|
||||
}, now); err != nil {
|
||||
t.Fatalf("upsert ssh config 4: %v", err)
|
||||
}
|
||||
if err := r.UpsertNodeSSHConfig(2, NftSSHConfigInput{
|
||||
Host: "203.0.113.2",
|
||||
Port: 22,
|
||||
Username: "admin",
|
||||
AuthType: "private_key",
|
||||
SudoMode: "sudo",
|
||||
}, now); err != nil {
|
||||
t.Fatalf("upsert ssh config 2: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNftablesNodesForCollection()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNftablesNodesForCollection: %v", err)
|
||||
}
|
||||
if len(nodes) != 2 {
|
||||
t.Fatalf("expected 2 collection nodes, got %d: %+v", len(nodes), nodes)
|
||||
}
|
||||
if nodes[0].NodeID != 2 || nodes[1].NodeID != 4 {
|
||||
t.Fatalf("expected nodes ordered by id [2 4], got [%d %d]", nodes[0].NodeID, nodes[1].NodeID)
|
||||
}
|
||||
if nodes[0].Config.NodeID != 2 || nodes[0].Config.Host != "203.0.113.2" || nodes[0].Config.Username != "admin" {
|
||||
t.Fatalf("unexpected first config: %+v", nodes[0].Config)
|
||||
}
|
||||
if nodes[1].Config.NodeID != 4 || nodes[1].Config.Port != 2222 || nodes[1].Config.Password.String != "secret-4" {
|
||||
t.Fatalf("unexpected second config: %+v", nodes[1].Config)
|
||||
}
|
||||
}
|
||||
|
||||
func seedCollectionNode(t *testing.T, r *Repository, id int64, forwardMode string, status int, now int64) {
|
||||
t.Helper()
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO node(id, name, secret, server_ip, port, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, forward_mode)
|
||||
VALUES(?, ?, 'secret', ?, '1000-2000', ?, ?, ?, '[::]', '[::]', 0, ?)
|
||||
`, id, "node", "198.51.100.1", now, now, status, forwardMode).Error; err != nil {
|
||||
t.Fatalf("insert node %d: %v", id, err)
|
||||
}
|
||||
}
|
||||
@@ -207,20 +207,23 @@ func (r *Repository) ListActiveForwardsByNode(nodeID int64) ([]model.ForwardReco
|
||||
}
|
||||
rows := make([]model.ForwardRecord, 0, len(forwards))
|
||||
for _, f := range forwards {
|
||||
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
|
||||
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,
|
||||
MaxConn: f.MaxConn,
|
||||
IPMaxConn: f.IPMaxConn,
|
||||
IPSpeedID: f.IPSpeedID,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
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,
|
||||
ProxyProtocolReceive: proxyProtocolReceive,
|
||||
ProxyProtocolSend: proxyProtocolSend,
|
||||
})
|
||||
}
|
||||
for i := range rows {
|
||||
@@ -237,3 +240,10 @@ func defaultString(value, fallback string) string {
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
|
||||
if send == 0 && legacy > 0 {
|
||||
send = legacy
|
||||
}
|
||||
return receive, send
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -123,6 +124,114 @@ func TestNftablesNodeModeSSHConfigAndBindingPersistence(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestNftablesNodeSSHConfigSurvivesRepositoryReopen(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg NftSSHConfigInput
|
||||
}{
|
||||
{
|
||||
name: "password",
|
||||
cfg: NftSSHConfigInput{
|
||||
Host: "203.0.113.10",
|
||||
Port: 2222,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "ssh-password",
|
||||
Passphrase: "key-passphrase",
|
||||
SudoMode: "sudo",
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "private-key",
|
||||
cfg: NftSSHConfigInput{
|
||||
Host: "203.0.113.11",
|
||||
Port: 2223,
|
||||
Username: "admin",
|
||||
AuthType: "private_key",
|
||||
PrivateKey: "PRIVATE-KEY-SHOULD-PERSIST",
|
||||
Passphrase: "key-passphrase",
|
||||
SudoMode: "sudo_su",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "nftables-ssh.sqlite")
|
||||
r, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.CreateNode(
|
||||
"nft-node",
|
||||
"secret",
|
||||
tt.cfg.Host,
|
||||
nil,
|
||||
nil,
|
||||
"10000-20000",
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
now,
|
||||
1,
|
||||
"[::]",
|
||||
"[::]",
|
||||
1,
|
||||
0,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
"nftables",
|
||||
); err != nil {
|
||||
t.Fatalf("CreateNode: %v", err)
|
||||
}
|
||||
|
||||
nodes, err := r.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("ListNodes: %v", err)
|
||||
}
|
||||
nodeID := nodes[0]["id"].(int64)
|
||||
if err := r.UpsertNodeSSHConfig(nodeID, tt.cfg, now); err != nil {
|
||||
t.Fatalf("UpsertNodeSSHConfig: %v", err)
|
||||
}
|
||||
if err := r.Close(); err != nil {
|
||||
t.Fatalf("close repo: %v", err)
|
||||
}
|
||||
|
||||
reopened, err := Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("reopen repo: %v", err)
|
||||
}
|
||||
defer reopened.Close()
|
||||
|
||||
loaded, err := reopened.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetNodeSSHConfig after reopen: %v", err)
|
||||
}
|
||||
if loaded.Host != tt.cfg.Host || loaded.Port != tt.cfg.Port || loaded.Username != tt.cfg.Username || loaded.AuthType != tt.cfg.AuthType || loaded.SudoMode != tt.cfg.SudoMode {
|
||||
t.Fatalf("unexpected ssh config after reopen: %+v", loaded)
|
||||
}
|
||||
if tt.cfg.Password != "" && (!loaded.Password.Valid || loaded.Password.String != tt.cfg.Password) {
|
||||
t.Fatalf("expected password to persist after reopen, got %+v", loaded.Password)
|
||||
}
|
||||
if tt.cfg.PrivateKey != "" && (!loaded.PrivateKey.Valid || loaded.PrivateKey.String != tt.cfg.PrivateKey) {
|
||||
t.Fatalf("expected private key to persist after reopen, got %+v", loaded.PrivateKey)
|
||||
}
|
||||
if tt.cfg.Passphrase != "" && (!loaded.Passphrase.Valid || loaded.Passphrase.String != tt.cfg.Passphrase) {
|
||||
t.Fatalf("expected passphrase to persist after reopen, got %+v", loaded.Passphrase)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateNodeWithoutForwardModePreservesExistingMode(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestTrafficLimitMiBSurvivesBackupRestore(t *testing.T) {
|
||||
source, err := Open(filepath.Join(t.TempDir(), "source.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer source.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
userID, err := source.CreateUser("mib-user", "hash", 1, now+86400000, 1, 1, 10, 1, 0, now, 500)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tunnel := model.Tunnel{Name: "mib-tunnel", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: now, UpdatedTime: now, Status: 1, Inx: 1}
|
||||
if err := source.DB().Create(&tunnel).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, _, err := source.EnsureUserTunnelGrant(userID, tunnel.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
grants, err := source.GetUserPackageTunnels(userID)
|
||||
if err != nil || len(grants) != 1 || grants[0].FlowMiB != 500 {
|
||||
t.Fatalf("inherited tunnel quota = %+v, err = %v", grants, err)
|
||||
}
|
||||
backup, err := source.ExportAll()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
found := false
|
||||
for _, user := range backup.Users {
|
||||
if user.User == "mib-user" {
|
||||
found = user.Flow == 1 && user.FlowMiB == 500
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("500 MiB user quota missing from backup")
|
||||
}
|
||||
if len(backup.UserTunnels) != 1 || backup.UserTunnels[0].FlowMiB != 500 {
|
||||
t.Fatalf("tunnel quota missing from backup: %+v", backup.UserTunnels)
|
||||
}
|
||||
|
||||
dest, err := Open(filepath.Join(t.TempDir(), "dest.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer dest.Close()
|
||||
if _, err := dest.Import(backup, []string{"users", "tunnels", "userTunnels"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
user, err := dest.GetUserByUsername("mib-user")
|
||||
if err != nil || user == nil || user.Flow != 1 || user.FlowMiB != 500 {
|
||||
t.Fatalf("restored quota = %+v, err = %v", user, err)
|
||||
}
|
||||
grants, err = dest.GetUserPackageTunnels(user.ID)
|
||||
if err != nil || len(grants) != 1 || grants[0].FlowMiB != 500 {
|
||||
t.Fatalf("restored tunnel quota = %+v, err = %v", grants, err)
|
||||
}
|
||||
}
|
||||
@@ -44,7 +44,8 @@ func (r *Repository) GetLatestTunnelQualities() ([]model.TunnelQuality, error) {
|
||||
// Use window function (works on modern SQLite 3.25+ and PostgreSQL).
|
||||
q := `
|
||||
SELECT id, tunnel_id, entry_to_exit_latency, exit_to_bing_latency,
|
||||
entry_to_exit_loss, exit_to_bing_loss, success, error_message, timestamp
|
||||
entry_to_exit_loss, exit_to_bing_loss, success, error_message, timestamp,
|
||||
chain_details
|
||||
FROM (
|
||||
SELECT *, ROW_NUMBER() OVER (PARTITION BY tunnel_id ORDER BY timestamp DESC, id DESC) AS rn
|
||||
FROM tunnel_quality
|
||||
|
||||
@@ -264,39 +264,11 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
var result map[int64]*model.UserQuotaView
|
||||
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
|
||||
var err error
|
||||
result, err = r.addUserQuotaUsageBatchTx(tx, usages, now)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -304,6 +276,48 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *Repository) addUserQuotaUsageBatchTx(tx *gorm.DB, usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
|
||||
if tx == nil {
|
||||
return nil, errors.New("database unavailable")
|
||||
}
|
||||
if len(usages) == 0 {
|
||||
return map[int64]*model.UserQuotaView{}, nil
|
||||
}
|
||||
|
||||
result := make(map[int64]*model.UserQuotaView, len(usages))
|
||||
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 nil, 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 nil, err
|
||||
}
|
||||
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
|
||||
}
|
||||
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")
|
||||
|
||||
@@ -2,6 +2,8 @@ package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
"github.com/go-gost/core/hop"
|
||||
@@ -38,11 +40,12 @@ type chainNamer interface {
|
||||
}
|
||||
|
||||
type Chain struct {
|
||||
name string
|
||||
hops []hop.Hop
|
||||
marker selector.Marker
|
||||
metadata metadata.Metadata
|
||||
logger logger.Logger
|
||||
name string
|
||||
hops []hop.Hop
|
||||
ownedHops []hop.Hop
|
||||
marker selector.Marker
|
||||
metadata metadata.Metadata
|
||||
logger logger.Logger
|
||||
}
|
||||
|
||||
func NewChain(name string, opts ...ChainOption) *Chain {
|
||||
@@ -61,8 +64,15 @@ func NewChain(name string, opts ...ChainOption) *Chain {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Chain) AddHop(hop hop.Hop) {
|
||||
func (c *Chain) AddHop(hop hop.Hop, owned ...bool) {
|
||||
c.hops = append(c.hops, hop)
|
||||
isOwned := true
|
||||
if len(owned) > 0 {
|
||||
isOwned = owned[0]
|
||||
}
|
||||
if isOwned {
|
||||
c.ownedHops = append(c.ownedHops, hop)
|
||||
}
|
||||
}
|
||||
|
||||
// Metadata implements metadata.Metadatable interface.
|
||||
@@ -112,6 +122,36 @@ func (c *Chain) Route(ctx context.Context, network, address string, opts ...chai
|
||||
return rt
|
||||
}
|
||||
|
||||
// Retire gracefully drains resources owned by a chain that has been replaced.
|
||||
func (c *Chain) Retire() {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
for _, h := range c.ownedHops {
|
||||
if retirer, ok := h.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
continue
|
||||
}
|
||||
if closer, ok := h.(io.Closer); ok {
|
||||
_ = closer.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases all resources owned by the chain.
|
||||
func (c *Chain) Close() error {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
for _, h := range c.ownedHops {
|
||||
if closer, ok := h.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
type chainGroup struct {
|
||||
chains []chain.Chainer
|
||||
selector selector.Selector[chain.Chainer]
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
corechain "github.com/go-gost/core/chain"
|
||||
corehop "github.com/go-gost/core/hop"
|
||||
)
|
||||
|
||||
type lifecycleTestHop struct {
|
||||
selected int
|
||||
retired int
|
||||
closed int
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Select(context.Context, ...corehop.SelectOption) *corechain.Node {
|
||||
h.selected++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Retire() {
|
||||
h.retired++
|
||||
}
|
||||
|
||||
func (h *lifecycleTestHop) Close() error {
|
||||
h.closed++
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestChainRoutesThroughSharedHopWithoutOwningLifecycle(t *testing.T) {
|
||||
hop := &lifecycleTestHop{}
|
||||
chain := NewChain("shared-hop")
|
||||
chain.AddHop(hop, false)
|
||||
|
||||
if route := chain.Route(context.Background(), "tcp", "example.com:443"); route == nil {
|
||||
t.Fatal("route is nil")
|
||||
}
|
||||
if hop.selected != 1 {
|
||||
t.Fatalf("shared hop selected %d times, want 1", hop.selected)
|
||||
}
|
||||
|
||||
chain.Retire()
|
||||
if err := chain.Close(); err != nil {
|
||||
t.Fatalf("close chain: %v", err)
|
||||
}
|
||||
if hop.retired != 0 || hop.closed != 0 {
|
||||
t.Fatalf("shared hop lifecycle changed: retired=%d closed=%d", hop.retired, hop.closed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainRetiresAndClosesOwnedHop(t *testing.T) {
|
||||
hop := &lifecycleTestHop{}
|
||||
chain := NewChain("owned-hop")
|
||||
chain.AddHop(hop)
|
||||
|
||||
chain.Retire()
|
||||
if err := chain.Close(); err != nil {
|
||||
t.Fatalf("close chain: %v", err)
|
||||
}
|
||||
if hop.retired != 1 || hop.closed != 1 {
|
||||
t.Fatalf("owned hop lifecycle: retired=%d closed=%d, want 1/1", hop.retired, hop.closed)
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,8 @@ package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
|
||||
"github.com/go-gost/core/chain"
|
||||
@@ -102,5 +104,34 @@ func (tr *Transport) Options() *chain.TransportOptions {
|
||||
func (tr *Transport) Copy() chain.Transporter {
|
||||
tr2 := &Transport{}
|
||||
*tr2 = *tr
|
||||
return tr
|
||||
return tr2
|
||||
}
|
||||
|
||||
// Retire prevents long-lived dialer sessions owned by an obsolete chain from
|
||||
// accepting new streams while allowing existing streams to drain.
|
||||
func (tr *Transport) Retire() {
|
||||
if tr == nil {
|
||||
return
|
||||
}
|
||||
if retirer, ok := tr.dialer.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
if retirer, ok := tr.connector.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases transport-owned dialer and connector resources.
|
||||
func (tr *Transport) Close() error {
|
||||
if tr == nil {
|
||||
return nil
|
||||
}
|
||||
var errs []error
|
||||
if closer, ok := tr.dialer.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
if closer, ok := tr.connector.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
package chain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
corechain "github.com/go-gost/core/chain"
|
||||
)
|
||||
|
||||
type copyTestRoute struct{}
|
||||
|
||||
func (copyTestRoute) Dial(context.Context, string, string, ...corechain.DialOption) (net.Conn, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (copyTestRoute) Bind(context.Context, string, string, ...corechain.BindOption) (net.Listener, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (copyTestRoute) Nodes() []*corechain.Node {
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestTransportCopyReturnsIndependentTransport(t *testing.T) {
|
||||
originalRoute := copyTestRoute{}
|
||||
replacementRoute := ©TestRoute{}
|
||||
original := NewTransport(nil, nil, corechain.RouteTransportOption(originalRoute))
|
||||
|
||||
copied, ok := original.Copy().(*Transport)
|
||||
if !ok {
|
||||
t.Fatalf("copy type = %T, want *Transport", original.Copy())
|
||||
}
|
||||
if copied == original {
|
||||
t.Fatal("Copy returned the original transport")
|
||||
}
|
||||
|
||||
copied.Options().Route = replacementRoute
|
||||
if original.Options().Route != originalRoute {
|
||||
t.Fatal("mutating copied transport changed original route")
|
||||
}
|
||||
if copied.Options().Route != replacementRoute {
|
||||
t.Fatal("copied transport did not retain its independent route")
|
||||
}
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package config
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"reflect"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -30,9 +31,87 @@ func Global() *Config {
|
||||
globalMux.RLock()
|
||||
defer globalMux.RUnlock()
|
||||
|
||||
cfg := &Config{}
|
||||
*cfg = *global
|
||||
return cfg
|
||||
return cloneConfig(global)
|
||||
}
|
||||
|
||||
// cloneConfig returns a detached snapshot of the runtime config. Runtime
|
||||
// commands mutate slices, pointers and metadata maps, so a shallow struct copy
|
||||
// is not sufficient once readers and persistence run concurrently.
|
||||
func cloneConfig(c *Config) *Config {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
v := cloneConfigValue(reflect.ValueOf(c))
|
||||
if !v.IsValid() || v.IsNil() {
|
||||
return nil
|
||||
}
|
||||
return v.Interface().(*Config)
|
||||
}
|
||||
|
||||
func cloneConfigValue(v reflect.Value) reflect.Value {
|
||||
if !v.IsValid() {
|
||||
return reflect.Value{}
|
||||
}
|
||||
|
||||
switch v.Kind() {
|
||||
case reflect.Interface:
|
||||
if v.IsNil() {
|
||||
return reflect.Zero(v.Type())
|
||||
}
|
||||
cloned := cloneConfigValue(v.Elem())
|
||||
out := reflect.New(v.Type()).Elem()
|
||||
if cloned.IsValid() && cloned.Type().AssignableTo(v.Type()) {
|
||||
out.Set(cloned)
|
||||
} else if cloned.IsValid() && cloned.Type().Implements(v.Type()) {
|
||||
out.Set(cloned)
|
||||
} else if cloned.IsValid() {
|
||||
out.Set(cloned)
|
||||
}
|
||||
return out
|
||||
case reflect.Pointer:
|
||||
if v.IsNil() {
|
||||
return reflect.Zero(v.Type())
|
||||
}
|
||||
out := reflect.New(v.Type().Elem())
|
||||
out.Elem().Set(cloneConfigValue(v.Elem()))
|
||||
return out
|
||||
case reflect.Map:
|
||||
if v.IsNil() {
|
||||
return reflect.Zero(v.Type())
|
||||
}
|
||||
out := reflect.MakeMapWithSize(v.Type(), v.Len())
|
||||
iter := v.MapRange()
|
||||
for iter.Next() {
|
||||
out.SetMapIndex(cloneConfigValue(iter.Key()), cloneConfigValue(iter.Value()))
|
||||
}
|
||||
return out
|
||||
case reflect.Slice:
|
||||
if v.IsNil() {
|
||||
return reflect.Zero(v.Type())
|
||||
}
|
||||
out := reflect.MakeSlice(v.Type(), v.Len(), v.Len())
|
||||
for i := 0; i < v.Len(); i++ {
|
||||
out.Index(i).Set(cloneConfigValue(v.Index(i)))
|
||||
}
|
||||
return out
|
||||
case reflect.Array:
|
||||
out := reflect.New(v.Type()).Elem()
|
||||
for i := 0; i < v.Len(); i++ {
|
||||
out.Index(i).Set(cloneConfigValue(v.Index(i)))
|
||||
}
|
||||
return out
|
||||
case reflect.Struct:
|
||||
out := reflect.New(v.Type()).Elem()
|
||||
out.Set(v)
|
||||
for i := 0; i < v.NumField(); i++ {
|
||||
if out.Field(i).CanSet() && v.Field(i).CanInterface() {
|
||||
out.Field(i).Set(cloneConfigValue(v.Field(i)))
|
||||
}
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return v
|
||||
}
|
||||
}
|
||||
|
||||
func Set(c *Config) {
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGlobalReturnsDetachedSnapshot(t *testing.T) {
|
||||
original := Global()
|
||||
t.Cleanup(func() { Set(original) })
|
||||
|
||||
Set(&Config{Services: []*ServiceConfig{{
|
||||
Name: "snapshot-service",
|
||||
Metadata: map[string]any{"paused": false},
|
||||
Handler: &HandlerConfig{
|
||||
Type: "relay",
|
||||
Metadata: map[string]any{"retries": 2},
|
||||
},
|
||||
}}})
|
||||
|
||||
snapshot := Global()
|
||||
snapshot.Services[0].Name = "changed"
|
||||
snapshot.Services[0].Metadata["paused"] = true
|
||||
snapshot.Services[0].Handler.Metadata["retries"] = 9
|
||||
|
||||
current := Global()
|
||||
if current.Services[0].Name != "snapshot-service" {
|
||||
t.Fatalf("snapshot mutated global service name: %q", current.Services[0].Name)
|
||||
}
|
||||
if paused, _ := current.Services[0].Metadata["paused"].(bool); paused {
|
||||
t.Fatalf("snapshot mutated global service metadata")
|
||||
}
|
||||
if retries, _ := current.Services[0].Handler.Metadata["retries"].(int); retries != 2 {
|
||||
t.Fatalf("snapshot mutated nested handler metadata: %v", current.Services[0].Handler.Metadata["retries"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentGlobalSnapshotAndUpdate(t *testing.T) {
|
||||
original := Global()
|
||||
t.Cleanup(func() { Set(original) })
|
||||
|
||||
Set(&Config{Services: []*ServiceConfig{{
|
||||
Name: "concurrent-service",
|
||||
Metadata: map[string]any{"generation": 0},
|
||||
}}})
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for worker := 0; worker < 4; worker++ {
|
||||
wg.Add(1)
|
||||
go func(worker int) {
|
||||
defer wg.Done()
|
||||
for i := 0; i < 500; i++ {
|
||||
if err := OnUpdate(func(c *Config) error {
|
||||
c.Services[0] = &ServiceConfig{
|
||||
Name: "concurrent-service",
|
||||
Metadata: map[string]any{"generation": worker*500 + i},
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Errorf("OnUpdate: %v", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}(worker)
|
||||
}
|
||||
|
||||
for i := 0; i < 2000; i++ {
|
||||
if _, err := json.Marshal(Global()); err != nil {
|
||||
t.Fatalf("marshal snapshot: %v", err)
|
||||
}
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestConcurrentPersistProducesValidConfig(t *testing.T) {
|
||||
original := Global()
|
||||
originalPath := PersistPath()
|
||||
persistMu.Lock()
|
||||
originalEnabled := persistEnable
|
||||
persistMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
Set(original)
|
||||
SetPersistPath(originalPath)
|
||||
persistMu.Lock()
|
||||
persistEnable = originalEnabled
|
||||
persistMu.Unlock()
|
||||
})
|
||||
|
||||
path := filepath.Join(t.TempDir(), "gost.json")
|
||||
SetPersistPath(path)
|
||||
EnablePersist()
|
||||
Set(&Config{Services: []*ServiceConfig{{Name: "persist-service"}}})
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for worker := 0; worker < 4; worker++ {
|
||||
wg.Add(1)
|
||||
go func(worker int) {
|
||||
defer wg.Done()
|
||||
for i := 0; i < 50; i++ {
|
||||
if err := OnUpdate(func(c *Config) error {
|
||||
c.Services[0].Metadata = map[string]any{"generation": worker*50 + i}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Errorf("persist update: %v", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}(worker)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
var persisted Config
|
||||
if err := persisted.ReadFile(path); err != nil {
|
||||
t.Fatalf("read persisted config: %v", err)
|
||||
}
|
||||
if len(persisted.Services) != 1 || persisted.Services[0] == nil || persisted.Services[0].Name != "persist-service" {
|
||||
t.Fatalf("unexpected persisted services: %#v", persisted.Services)
|
||||
}
|
||||
}
|
||||
@@ -35,16 +35,18 @@ func ParseChain(cfg *config.ChainConfig, log logger.Logger) (chain.Chainer, erro
|
||||
for _, ch := range cfg.Hops {
|
||||
var hop hop.Hop
|
||||
var err error
|
||||
owned := false
|
||||
|
||||
if ch.Nodes != nil || ch.Plugin != nil {
|
||||
if hop, err = hop_parser.ParseHop(ch, log); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
owned = true
|
||||
} else {
|
||||
hop = registry.HopRegistry().Get(ch.Name)
|
||||
}
|
||||
if hop != nil {
|
||||
c.AddHop(hop)
|
||||
c.AddHop(hop, owned)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -42,9 +42,9 @@ func EnablePersist() {
|
||||
// persist writes the current global config to the configured file atomically.
|
||||
func persist() error {
|
||||
persistMu.Lock()
|
||||
defer persistMu.Unlock()
|
||||
path := persistPath
|
||||
enabled := persistEnable
|
||||
persistMu.Unlock()
|
||||
|
||||
if !enabled || path == "" {
|
||||
return nil
|
||||
|
||||
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session.session == nil {
|
||||
if session == nil || session.session == nil {
|
||||
return true
|
||||
}
|
||||
return session.session.IsClosed()
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
kcp_util "github.com/go-gost/x/internal/util/kcp"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
mdutil "github.com/go-gost/x/metadata/util"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/xtaci/kcp-go/v5"
|
||||
@@ -25,6 +26,7 @@ func init() {
|
||||
type kcpDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -64,6 +66,9 @@ func (d *kcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOp
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -171,3 +176,36 @@ func (d *kcpDialer) initSession(ctx context.Context, addr net.Addr, conn net.Pac
|
||||
func (d *kcpDialer) Multiplex() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// Retire drains existing streams and closes their backing sessions once idle.
|
||||
func (d *kcpDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
// Close immediately releases all cached multiplex sessions.
|
||||
func (d *kcpDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *kcpDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
package kcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*kcpDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
@@ -21,6 +22,7 @@ func init() {
|
||||
type mtcpDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -55,6 +57,9 @@ func (d *mtcpDialer) Multiplex() bool {
|
||||
func (d *mtcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -88,6 +93,10 @@ func (d *mtcpDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
if d.md.handshakeTimeout > 0 {
|
||||
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
|
||||
@@ -129,3 +138,34 @@ func (d *mtcpDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
|
||||
}
|
||||
return &muxSession{conn: conn, session: session}, nil
|
||||
}
|
||||
|
||||
func (d *mtcpDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mtcpDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *mtcpDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package mtcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*mtcpDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
|
||||
conn, peer := net.Pipe()
|
||||
defer peer.Close()
|
||||
session := &muxSession{conn: conn}
|
||||
|
||||
if err := session.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("pre-handshake session still reports open after Close")
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
@@ -22,6 +23,7 @@ func init() {
|
||||
type mtlsDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
logger logger.Logger
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -56,6 +58,9 @@ func (d *mtlsDialer) Multiplex() bool {
|
||||
func (d *mtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -89,6 +94,10 @@ func (d *mtlsDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
if d.md.handshakeTimeout > 0 {
|
||||
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
|
||||
@@ -136,3 +145,34 @@ func (d *mtlsDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
|
||||
}
|
||||
return &muxSession{conn: conn, session: session}, nil
|
||||
}
|
||||
|
||||
func (d *mtlsDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mtlsDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *mtlsDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package mtls
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*mtlsDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
|
||||
conn, peer := net.Pipe()
|
||||
defer peer.Close()
|
||||
session := &muxSession{conn: conn}
|
||||
|
||||
if err := session.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("pre-handshake session still reports open after Close")
|
||||
}
|
||||
}
|
||||
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
|
||||
}
|
||||
|
||||
func (session *muxSession) Close() error {
|
||||
if session == nil {
|
||||
return nil
|
||||
}
|
||||
if session.session == nil {
|
||||
if session.conn != nil {
|
||||
conn := session.conn
|
||||
session.conn = nil
|
||||
return conn.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return session.session.Close()
|
||||
}
|
||||
|
||||
func (session *muxSession) IsClosed() bool {
|
||||
if session == nil {
|
||||
return true
|
||||
}
|
||||
if session.session == nil {
|
||||
return true
|
||||
}
|
||||
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
|
||||
}
|
||||
|
||||
func (session *muxSession) NumStreams() int {
|
||||
if session == nil || session.session == nil {
|
||||
return 0
|
||||
}
|
||||
return session.session.NumStreams()
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/go-gost/core/logger"
|
||||
md "github.com/go-gost/core/metadata"
|
||||
"github.com/go-gost/x/internal/util/mux"
|
||||
"github.com/go-gost/x/internal/util/sessionretire"
|
||||
ws_util "github.com/go-gost/x/internal/util/ws"
|
||||
"github.com/go-gost/x/registry"
|
||||
"github.com/gorilla/websocket"
|
||||
@@ -25,6 +26,7 @@ func init() {
|
||||
type mwsDialer struct {
|
||||
sessions map[string]*muxSession
|
||||
sessionMutex sync.Mutex
|
||||
retired bool
|
||||
tlsEnabled bool
|
||||
md metadata
|
||||
options dialer.Options
|
||||
@@ -70,6 +72,9 @@ func (d *mwsDialer) Multiplex() bool {
|
||||
func (d *mwsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[addr]
|
||||
if session != nil && session.IsClosed() {
|
||||
@@ -108,6 +113,10 @@ func (d *mwsDialer) Handshake(ctx context.Context, conn net.Conn, options ...dia
|
||||
|
||||
d.sessionMutex.Lock()
|
||||
defer d.sessionMutex.Unlock()
|
||||
if d.retired {
|
||||
conn.Close()
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
|
||||
session, ok := d.sessions[opts.Addr]
|
||||
if session != nil && session.conn != conn {
|
||||
@@ -208,3 +217,34 @@ func (d *mwsDialer) keepAlive(conn ws_util.WebsocketConn) {
|
||||
conn.SetWriteDeadline(time.Time{})
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mwsDialer) Retire() {
|
||||
for _, session := range d.detachSessions() {
|
||||
sessionretire.Gracefully(session)
|
||||
}
|
||||
}
|
||||
|
||||
func (d *mwsDialer) Close() error {
|
||||
var errs []error
|
||||
for _, session := range d.detachSessions() {
|
||||
errs = append(errs, session.Close())
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
func (d *mwsDialer) detachSessions() []*muxSession {
|
||||
if d == nil {
|
||||
return nil
|
||||
}
|
||||
d.sessionMutex.Lock()
|
||||
d.retired = true
|
||||
sessions := make([]*muxSession, 0, len(d.sessions))
|
||||
for _, session := range d.sessions {
|
||||
if session != nil {
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
}
|
||||
d.sessions = make(map[string]*muxSession)
|
||||
d.sessionMutex.Unlock()
|
||||
return sessions
|
||||
}
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package mws
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
|
||||
dialer := NewDialer().(*mwsDialer)
|
||||
dialer.Retire()
|
||||
|
||||
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
|
||||
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
|
||||
conn, peer := net.Pipe()
|
||||
defer peer.Close()
|
||||
session := &muxSession{conn: conn}
|
||||
|
||||
if err := session.Close(); err != nil {
|
||||
t.Fatalf("Close: %v", err)
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("pre-handshake session still reports open after Close")
|
||||
}
|
||||
}
|
||||
+49
-8
@@ -3,6 +3,7 @@ package hop
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sort"
|
||||
@@ -92,6 +93,7 @@ type chainHop struct {
|
||||
nodes []*chain.Node
|
||||
mu sync.RWMutex
|
||||
cancelFunc context.CancelFunc
|
||||
stopOnce sync.Once
|
||||
options options
|
||||
}
|
||||
|
||||
@@ -383,13 +385,52 @@ func (p *chainHop) parseNode(r io.Reader) ([]*chain.Node, error) {
|
||||
return nodes, nil
|
||||
}
|
||||
|
||||
func (p *chainHop) Close() error {
|
||||
p.cancelFunc()
|
||||
if p.options.fileLoader != nil {
|
||||
p.options.fileLoader.Close()
|
||||
func (p *chainHop) stopReload() {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
if p.options.redisLoader != nil {
|
||||
p.options.redisLoader.Close()
|
||||
}
|
||||
return nil
|
||||
p.stopOnce.Do(func() {
|
||||
p.cancelFunc()
|
||||
if p.options.fileLoader != nil {
|
||||
p.options.fileLoader.Close()
|
||||
}
|
||||
if p.options.redisLoader != nil {
|
||||
p.options.redisLoader.Close()
|
||||
}
|
||||
if p.options.httpLoader != nil {
|
||||
p.options.httpLoader.Close()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func (p *chainHop) Retire() {
|
||||
if p == nil {
|
||||
return
|
||||
}
|
||||
p.stopReload()
|
||||
for _, node := range p.Nodes() {
|
||||
if node == nil || node.Options().Transport == nil {
|
||||
continue
|
||||
}
|
||||
if retirer, ok := node.Options().Transport.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *chainHop) Close() error {
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
p.stopReload()
|
||||
var errs []error
|
||||
for _, node := range p.Nodes() {
|
||||
if node == nil || node.Options().Transport == nil {
|
||||
continue
|
||||
}
|
||||
if closer, ok := node.Options().Transport.(io.Closer); ok {
|
||||
errs = append(errs, closer.Close())
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
package sessionretire
|
||||
|
||||
import "time"
|
||||
|
||||
const (
|
||||
defaultIdleGrace = time.Second
|
||||
defaultPollPeriod = 100 * time.Millisecond
|
||||
)
|
||||
|
||||
// Session is the lifecycle surface shared by the multiplexed dialers.
|
||||
type Session interface {
|
||||
Close() error
|
||||
IsClosed() bool
|
||||
NumStreams() int
|
||||
}
|
||||
|
||||
// Gracefully closes a retired session after all existing streams have drained.
|
||||
// A short idle grace covers the Dial/Handshake hand-off used by several dialers.
|
||||
func Gracefully(session Session) {
|
||||
if session == nil {
|
||||
return
|
||||
}
|
||||
go waitUntilIdle(session, defaultIdleGrace, defaultPollPeriod)
|
||||
}
|
||||
|
||||
func waitUntilIdle(session Session, idleGrace, pollPeriod time.Duration) {
|
||||
if session == nil {
|
||||
return
|
||||
}
|
||||
if idleGrace <= 0 {
|
||||
idleGrace = defaultIdleGrace
|
||||
}
|
||||
if pollPeriod <= 0 {
|
||||
pollPeriod = defaultPollPeriod
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(pollPeriod)
|
||||
defer ticker.Stop()
|
||||
|
||||
var idleSince time.Time
|
||||
for {
|
||||
if session.IsClosed() {
|
||||
_ = session.Close()
|
||||
return
|
||||
}
|
||||
if session.NumStreams() == 0 {
|
||||
if idleSince.IsZero() {
|
||||
idleSince = time.Now()
|
||||
} else if time.Since(idleSince) >= idleGrace {
|
||||
_ = session.Close()
|
||||
return
|
||||
}
|
||||
} else {
|
||||
idleSince = time.Time{}
|
||||
}
|
||||
<-ticker.C
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package sessionretire
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
type testSession struct {
|
||||
mu sync.Mutex
|
||||
streams int
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (s *testSession) Close() error {
|
||||
s.mu.Lock()
|
||||
s.closed = true
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *testSession) IsClosed() bool {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.closed
|
||||
}
|
||||
|
||||
func (s *testSession) NumStreams() int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.streams
|
||||
}
|
||||
|
||||
func TestWaitUntilIdlePreservesActiveStreams(t *testing.T) {
|
||||
session := &testSession{streams: 1}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
waitUntilIdle(session, 20*time.Millisecond, time.Millisecond)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
time.Sleep(30 * time.Millisecond)
|
||||
if session.IsClosed() {
|
||||
t.Fatal("active session was closed")
|
||||
}
|
||||
|
||||
session.mu.Lock()
|
||||
session.streams = 0
|
||||
session.mu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(250 * time.Millisecond):
|
||||
t.Fatal("idle session was not closed")
|
||||
}
|
||||
if !session.IsClosed() {
|
||||
t.Fatal("retired session did not close after becoming idle")
|
||||
}
|
||||
}
|
||||
@@ -28,7 +28,13 @@ func (r *chainRegistry) Register(name string, v chain.Chainer) error {
|
||||
}
|
||||
|
||||
func (r *chainRegistry) replace(name string, v chain.Chainer) {
|
||||
r.m.Store(name, v)
|
||||
old, loaded := r.m.Swap(name, v)
|
||||
if !loaded {
|
||||
return
|
||||
}
|
||||
if retirer, ok := old.(interface{ Retire() }); ok {
|
||||
retirer.Retire()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *chainRegistry) Get(name string) chain.Chainer {
|
||||
|
||||
@@ -16,6 +16,15 @@ func (c testChainer) Route(context.Context, string, string, ...chain.RouteOption
|
||||
return c.route
|
||||
}
|
||||
|
||||
type retiringTestChainer struct {
|
||||
testChainer
|
||||
retired bool
|
||||
}
|
||||
|
||||
func (c *retiringTestChainer) Retire() {
|
||||
c.retired = true
|
||||
}
|
||||
|
||||
type testRoute struct {
|
||||
nodes []*chain.Node
|
||||
}
|
||||
@@ -49,3 +58,20 @@ func TestReplaceChainOverwritesExistingRegistration(t *testing.T) {
|
||||
t.Fatalf("expected replacement chain route, got %#v", route)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceChainRetiresPreviousRegistration(t *testing.T) {
|
||||
name := "replace_chain_retire_tdd"
|
||||
ChainRegistry().Unregister(name)
|
||||
defer ChainRegistry().Unregister(name)
|
||||
|
||||
old := &retiringTestChainer{}
|
||||
if err := ChainRegistry().Register(name, old); err != nil {
|
||||
t.Fatalf("register old chain: %v", err)
|
||||
}
|
||||
if err := ReplaceChain(name, testChainer{}); err != nil {
|
||||
t.Fatalf("replace chain: %v", err)
|
||||
}
|
||||
if !old.retired {
|
||||
t.Fatal("previous chain was not retired")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
package socket
|
||||
|
||||
import (
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
coreservice "github.com/go-gost/core/service"
|
||||
"github.com/go-gost/x/config"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
|
||||
type blockingCommandService struct {
|
||||
started chan struct{}
|
||||
release chan struct{}
|
||||
startedOnce sync.Once
|
||||
}
|
||||
|
||||
func (s *blockingCommandService) Serve() error { return nil }
|
||||
func (s *blockingCommandService) Addr() net.Addr { return nil }
|
||||
func (s *blockingCommandService) Close() error {
|
||||
s.startedOnce.Do(func() { close(s.started) })
|
||||
<-s.release
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestMutationCommandsAreSerialized(t *testing.T) {
|
||||
mutations := []string{
|
||||
"AddService", "UpdateService", "DeleteService", "PauseService", "ResumeService",
|
||||
"AddChains", "UpdateChains", "DeleteChains",
|
||||
"AddLimiters", "UpdateLimiters", "DeleteLimiters",
|
||||
"AddCLimiters", "UpdateCLimiters", "DeleteCLimiters",
|
||||
"SetProtocol", "UpgradeAgent", "RollbackAgent", "reload",
|
||||
}
|
||||
for _, command := range mutations {
|
||||
if !isMutationCommand(command) {
|
||||
t.Fatalf("expected %s to use the serialized mutation queue", command)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadOnlyCommandsRemainBoundedAsync(t *testing.T) {
|
||||
for _, command := range []string{"TcpPing", "UdpPing", "ServiceMonitorCheck"} {
|
||||
if isMutationCommand(command) {
|
||||
t.Fatalf("expected %s to remain a read-only command", command)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCommandResponseTypePreservesRequestContract(t *testing.T) {
|
||||
if got := commandResponseType("UpdateService"); got != "UpdateServiceResponse" {
|
||||
t.Fatalf("unexpected response type: %s", got)
|
||||
}
|
||||
if got := commandResponseType(""); got != "UnknownCommandResponse" {
|
||||
t.Fatalf("unexpected empty command response type: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) {
|
||||
originalConfig := config.Global()
|
||||
t.Cleanup(func() { config.Set(originalConfig) })
|
||||
|
||||
firstName := "mutation_queue_first_tdd"
|
||||
secondName := "mutation_queue_second_tdd"
|
||||
first := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
|
||||
second := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
|
||||
for _, name := range []string{firstName, secondName} {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
select {
|
||||
case <-first.release:
|
||||
default:
|
||||
close(first.release)
|
||||
}
|
||||
select {
|
||||
case <-second.release:
|
||||
default:
|
||||
close(second.release)
|
||||
}
|
||||
registry.ServiceRegistry().Unregister(firstName)
|
||||
registry.ServiceRegistry().Unregister(secondName)
|
||||
})
|
||||
if err := registry.ServiceRegistry().Register(firstName, coreservice.Service(first)); err != nil {
|
||||
t.Fatalf("register first service: %v", err)
|
||||
}
|
||||
if err := registry.ServiceRegistry().Register(secondName, coreservice.Service(second)); err != nil {
|
||||
t.Fatalf("register second service: %v", err)
|
||||
}
|
||||
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: firstName}, {Name: secondName}}})
|
||||
|
||||
reporter := NewWebSocketReporter("", "mutation-queue-test-secret")
|
||||
go reporter.runMutationCommands()
|
||||
t.Cleanup(reporter.Stop)
|
||||
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{firstName}}})
|
||||
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{secondName}}})
|
||||
|
||||
select {
|
||||
case <-first.started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("first mutation did not start")
|
||||
}
|
||||
select {
|
||||
case <-second.started:
|
||||
t.Fatal("second mutation started before first mutation completed")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
|
||||
close(first.release)
|
||||
select {
|
||||
case <-second.started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("second mutation did not start after first mutation completed")
|
||||
}
|
||||
close(second.release)
|
||||
}
|
||||
+82
-18
@@ -15,6 +15,11 @@ import (
|
||||
xservice "github.com/go-gost/x/service"
|
||||
)
|
||||
|
||||
type serviceReplacement struct {
|
||||
config config.ServiceConfig
|
||||
oldConfig *config.ServiceConfig
|
||||
}
|
||||
|
||||
func createServices(req createServicesRequest) error {
|
||||
|
||||
if len(req.Data) == 0 {
|
||||
@@ -95,11 +100,11 @@ func updateServices(req updateServicesRequest) error {
|
||||
req.Data[i].Name = name
|
||||
}
|
||||
|
||||
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)
|
||||
changedServices := make([]struct {
|
||||
config config.ServiceConfig
|
||||
service service.Service
|
||||
}, 0, len(req.Data))
|
||||
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)。
|
||||
// 配置变更命令由 WebSocket reporter 串行调度,但这里仍保留完整回滚,
|
||||
// 避免新配置解析或监听失败后旧服务永久消失。
|
||||
originalConfig := config.Global()
|
||||
changedServices := make([]serviceReplacement, 0, len(req.Data))
|
||||
for i := range req.Data {
|
||||
serviceConfig := &req.Data[i]
|
||||
name := serviceConfig.Name
|
||||
@@ -107,33 +112,48 @@ func updateServices(req updateServicesRequest) error {
|
||||
continue
|
||||
}
|
||||
|
||||
// 1. 获取旧服务
|
||||
old := registry.ServiceRegistry().Get(name)
|
||||
var oldConfig *config.ServiceConfig
|
||||
if originalConfig != nil {
|
||||
for _, current := range originalConfig.Services {
|
||||
if current != nil && strings.TrimSpace(current.Name) == name {
|
||||
oldConfig = current
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2. 关闭旧服务 (如果存在)
|
||||
if old != nil {
|
||||
// 3. 从注册表移除旧服务;registry 会负责关闭旧服务。
|
||||
// 1. 关闭并移除旧服务(如果存在)。同名监听必须先释放端口,
|
||||
// 才能创建新 listener。
|
||||
if registry.ServiceRegistry().Get(name) != nil {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
|
||||
// 4. 解析新服务配置
|
||||
// 2. 解析新服务配置。
|
||||
svc, err := parser.ParseService(serviceConfig)
|
||||
if err != nil {
|
||||
rollbackErr := restoreServiceRuntime(name, oldConfig)
|
||||
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
|
||||
if rollbackErr != nil {
|
||||
return fmt.Errorf("create service %s failed: %v; restore previous service failed: %w", name, err, rollbackErr)
|
||||
}
|
||||
return errors.New("create service " + name + " failed: " + err.Error())
|
||||
}
|
||||
changedServices = append(changedServices, struct {
|
||||
config config.ServiceConfig
|
||||
service service.Service
|
||||
}{*serviceConfig, svc})
|
||||
|
||||
// 5. 注册新服务
|
||||
// 3. 注册并启动新服务。
|
||||
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
||||
svc.Close()
|
||||
rollbackErr := restoreServiceRuntime(name, oldConfig)
|
||||
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
|
||||
if rollbackErr != nil {
|
||||
return fmt.Errorf("service %s already exists; restore previous service failed: %w", name, rollbackErr)
|
||||
}
|
||||
return errors.New("service " + name + " already exists")
|
||||
}
|
||||
|
||||
// 6. 启动新服务
|
||||
go svc.Serve()
|
||||
changedServices = append(changedServices, serviceReplacement{
|
||||
config: *serviceConfig,
|
||||
oldConfig: oldConfig,
|
||||
})
|
||||
}
|
||||
if len(changedServices) == 0 {
|
||||
return nil
|
||||
@@ -158,12 +178,56 @@ func updateServices(req updateServicesRequest) error {
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
config.Set(originalConfig)
|
||||
if rollbackErr := rollbackServiceReplacements(changedServices); rollbackErr != nil {
|
||||
return fmt.Errorf("%w; restore previous services failed: %v", err, rollbackErr)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func restoreServiceRuntime(name string, serviceConfig *config.ServiceConfig) error {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" || serviceConfig == nil {
|
||||
return nil
|
||||
}
|
||||
if registry.ServiceRegistry().Get(name) != nil {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
|
||||
cfgCopy := *serviceConfig
|
||||
cfgCopy.Name = name
|
||||
svc, err := parser.ParseService(&cfgCopy)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
|
||||
svc.Close()
|
||||
return err
|
||||
}
|
||||
go svc.Serve()
|
||||
return nil
|
||||
}
|
||||
|
||||
func rollbackServiceReplacements(replacements []serviceReplacement) error {
|
||||
var rollbackErr error
|
||||
for i := len(replacements) - 1; i >= 0; i-- {
|
||||
name := strings.TrimSpace(replacements[i].config.Name)
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
if registry.ServiceRegistry().Get(name) != nil {
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
}
|
||||
if err := restoreServiceRuntime(name, replacements[i].oldConfig); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, fmt.Errorf("restore service %s: %w", name, err))
|
||||
}
|
||||
}
|
||||
return rollbackErr
|
||||
}
|
||||
|
||||
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
|
||||
cfg := config.Global()
|
||||
if cfg == nil {
|
||||
|
||||
@@ -7,6 +7,8 @@ import (
|
||||
corelogger "github.com/go-gost/core/logger"
|
||||
"github.com/go-gost/core/service"
|
||||
"github.com/go-gost/x/config"
|
||||
_ "github.com/go-gost/x/handler/auto"
|
||||
_ "github.com/go-gost/x/listener/tcp"
|
||||
xlogger "github.com/go-gost/x/logger"
|
||||
"github.com/go-gost/x/registry"
|
||||
)
|
||||
@@ -15,6 +17,43 @@ type recordingService struct {
|
||||
closed int
|
||||
}
|
||||
|
||||
func TestUpdateServicesParseFailureRestoresPreviousRuntime(t *testing.T) {
|
||||
corelogger.SetDefault(xlogger.Nop())
|
||||
|
||||
name := "restore_after_failed_update_tdd"
|
||||
existing := &recordingService{}
|
||||
registry.ServiceRegistry().Unregister(name)
|
||||
t.Cleanup(func() { registry.ServiceRegistry().Unregister(name) })
|
||||
if err := registry.ServiceRegistry().Register(name, service.Service(existing)); err != nil {
|
||||
t.Fatalf("register existing service: %v", err)
|
||||
}
|
||||
|
||||
originalConfig := config.Global()
|
||||
t.Cleanup(func() { config.Set(originalConfig) })
|
||||
serviceConfig := config.ServiceConfig{Name: name, Addr: "127.0.0.1:0"}
|
||||
config.Set(&config.Config{Services: []*config.ServiceConfig{&serviceConfig}})
|
||||
|
||||
invalid := serviceConfig
|
||||
invalid.Listener = &config.ListenerConfig{Type: "listener-does-not-exist"}
|
||||
if err := updateServices(updateServicesRequest{Data: []config.ServiceConfig{invalid}}); err == nil {
|
||||
t.Fatalf("expected invalid service update to fail")
|
||||
}
|
||||
if existing.closed != 1 {
|
||||
t.Fatalf("expected old runtime to be closed once, got %d", existing.closed)
|
||||
}
|
||||
if registry.ServiceRegistry().Get(name) == nil {
|
||||
t.Fatalf("expected previous runtime to be restored")
|
||||
}
|
||||
|
||||
cfg := config.Global()
|
||||
if len(cfg.Services) != 1 || cfg.Services[0] == nil || cfg.Services[0].Name != name {
|
||||
t.Fatalf("expected previous config to remain, got %#v", cfg.Services)
|
||||
}
|
||||
if cfg.Services[0].Listener != nil {
|
||||
t.Fatalf("expected previous listener config to remain unchanged")
|
||||
}
|
||||
}
|
||||
|
||||
func (s *recordingService) Serve() error { return nil }
|
||||
func (s *recordingService) Addr() net.Addr { return nil }
|
||||
func (s *recordingService) Close() error {
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"runtime/debug"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -151,6 +152,9 @@ const (
|
||||
initialBackoff = 2 * time.Second // 重连初始退避
|
||||
maxBackoff = 2 * time.Minute // 重连最大退避
|
||||
defaultMetricReportInterval = 5 * time.Second
|
||||
maxConcurrentTCPPings = 8
|
||||
maxConcurrentReadCommands = 16
|
||||
maxQueuedMutationCommands = 256
|
||||
)
|
||||
|
||||
type WebSocketReporter struct {
|
||||
@@ -172,6 +176,9 @@ type WebSocketReporter struct {
|
||||
connecting bool // 正在连接状态
|
||||
connMutex sync.Mutex // 连接状态锁
|
||||
aesCrypto *crypto.AESCrypto // AES加密器
|
||||
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
|
||||
readCommandSem chan struct{} // 限制只读命令并发,避免诊断请求耗尽资源
|
||||
mutationQueue chan CommandMessage
|
||||
}
|
||||
|
||||
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
|
||||
@@ -201,11 +208,37 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
|
||||
connected: false,
|
||||
connecting: false,
|
||||
aesCrypto: aesCrypto,
|
||||
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
|
||||
readCommandSem: make(chan struct{}, maxConcurrentReadCommands),
|
||||
mutationQueue: make(chan CommandMessage, maxQueuedMutationCommands),
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebSocketReporter) tryAcquireTCPPingSlot() bool {
|
||||
if w == nil || w.tcpPingSem == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case w.tcpPingSem <- struct{}{}:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebSocketReporter) releaseTCPPingSlot() {
|
||||
if w == nil || w.tcpPingSem == nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-w.tcpPingSem:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// Start 启动WebSocket报告器
|
||||
func (w *WebSocketReporter) Start() {
|
||||
go w.runMutationCommands()
|
||||
go w.run()
|
||||
}
|
||||
|
||||
@@ -752,8 +785,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
||||
}
|
||||
|
||||
if cmdMsg.Type != "call" {
|
||||
// 所有命令统一异步执行,避免阻塞消息接收循环
|
||||
go w.routeCommand(cmdMsg)
|
||||
w.dispatchCommand(cmdMsg)
|
||||
}
|
||||
} else {
|
||||
// 处理普通消息
|
||||
@@ -764,8 +796,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
||||
return
|
||||
}
|
||||
if cmdMsg.Type != "call" {
|
||||
// 所有命令统一异步执行,避免阻塞消息接收循环
|
||||
go w.routeCommand(cmdMsg)
|
||||
w.dispatchCommand(cmdMsg)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -774,6 +805,86 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
|
||||
}
|
||||
}
|
||||
|
||||
// dispatchCommand keeps all runtime mutations ordered while allowing bounded
|
||||
// concurrency for read-only diagnostics. Mutations share process-wide
|
||||
// registries and configuration, so running them concurrently can corrupt the
|
||||
// persisted config or interleave service lifecycle operations.
|
||||
func (w *WebSocketReporter) dispatchCommand(cmd CommandMessage) {
|
||||
if isMutationCommand(cmd.Type) {
|
||||
select {
|
||||
case w.mutationQueue <- cmd:
|
||||
case <-w.ctx.Done():
|
||||
w.sendCommandFailure(cmd, "Agent is shutting down")
|
||||
default:
|
||||
w.sendCommandFailure(cmd, "运行时配置命令队列已满,请稍后重试")
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
select {
|
||||
case w.readCommandSem <- struct{}{}:
|
||||
go func() {
|
||||
defer func() { <-w.readCommandSem }()
|
||||
w.routeCommandSafely(cmd)
|
||||
}()
|
||||
case <-w.ctx.Done():
|
||||
w.sendCommandFailure(cmd, "Agent is shutting down")
|
||||
default:
|
||||
w.sendCommandFailure(cmd, "只读命令并发过多,请稍后重试")
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebSocketReporter) runMutationCommands() {
|
||||
for {
|
||||
select {
|
||||
case <-w.ctx.Done():
|
||||
return
|
||||
case cmd := <-w.mutationQueue:
|
||||
w.routeCommandSafely(cmd)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *WebSocketReporter) routeCommandSafely(cmd CommandMessage) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
fmt.Printf("❌ 命令处理 panic: type=%s panic=%v\n%s", cmd.Type, recovered, debug.Stack())
|
||||
w.sendCommandFailure(cmd, fmt.Sprintf("命令处理异常: %v", recovered))
|
||||
}
|
||||
}()
|
||||
w.routeCommand(cmd)
|
||||
}
|
||||
|
||||
func (w *WebSocketReporter) sendCommandFailure(cmd CommandMessage, message string) {
|
||||
w.sendResponse(CommandResponse{
|
||||
Type: commandResponseType(cmd.Type),
|
||||
Success: false,
|
||||
Message: message,
|
||||
RequestId: cmd.RequestId,
|
||||
})
|
||||
}
|
||||
|
||||
func commandResponseType(commandType string) string {
|
||||
commandType = strings.TrimSpace(commandType)
|
||||
if commandType == "" {
|
||||
return "UnknownCommandResponse"
|
||||
}
|
||||
return commandType + "Response"
|
||||
}
|
||||
|
||||
func isMutationCommand(commandType string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(commandType)) {
|
||||
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice",
|
||||
"addchains", "updatechains", "deletechains",
|
||||
"addlimiters", "updatelimiters", "deletelimiters",
|
||||
"addclimiters", "updateclimiters", "deleteclimiters",
|
||||
"setprotocol", "upgradeagent", "rollbackagent", "reload":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// routeCommand 路由命令到对应的处理函数
|
||||
func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
jsonBytes, errs := json.Marshal(cmd)
|
||||
@@ -840,9 +951,14 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
|
||||
|
||||
// TCP Ping 诊断命令(只读,不需要保存配置)
|
||||
case "TcpPing":
|
||||
response.Type = "TcpPingResponse"
|
||||
if !w.tryAcquireTCPPingSlot() {
|
||||
err = fmt.Errorf("TCP探测任务过多,请稍后重试")
|
||||
break
|
||||
}
|
||||
defer w.releaseTCPPingSlot()
|
||||
var tcpPingResult TcpPingResponse
|
||||
tcpPingResult, err = w.handleTcpPing(cmd.Data)
|
||||
response.Type = "TcpPingResponse"
|
||||
response.Data = tcpPingResult
|
||||
// needSaveConfig = false (默认值)
|
||||
|
||||
|
||||
@@ -148,6 +148,25 @@ func TestNewWebSocketReporterUsesReducedMetricInterval(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebSocketReporterLimitsConcurrentTCPPings(t *testing.T) {
|
||||
reporter := &WebSocketReporter{tcpPingSem: make(chan struct{}, maxConcurrentTCPPings)}
|
||||
for i := 0; i < maxConcurrentTCPPings; i++ {
|
||||
if !reporter.tryAcquireTCPPingSlot() {
|
||||
t.Fatalf("expected TCP ping slot %d to be available", i)
|
||||
}
|
||||
}
|
||||
if reporter.tryAcquireTCPPingSlot() {
|
||||
t.Fatalf("expected TCP ping concurrency limit at %d", maxConcurrentTCPPings)
|
||||
}
|
||||
for i := 0; i < maxConcurrentTCPPings; i++ {
|
||||
reporter.releaseTCPPingSlot()
|
||||
}
|
||||
if !reporter.tryAcquireTCPPingSlot() {
|
||||
t.Fatalf("expected released TCP ping slot to be reusable")
|
||||
}
|
||||
reporter.releaseTCPPingSlot()
|
||||
}
|
||||
|
||||
func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
|
||||
err := errors.New("websocket: bad handshake")
|
||||
resp := &http.Response{
|
||||
|
||||
+290
-42
@@ -1,4 +1,31 @@
|
||||
#!/bin/bash
|
||||
#!/bin/sh
|
||||
# shellcheck shell=bash
|
||||
|
||||
# Alpine 默认不带 Bash。先用系统自带的 /bin/sh 安装/切换到 Bash,
|
||||
# 后续主体继续使用 Bash 语法,避免要求用户手动准备运行环境。
|
||||
if [ -z "${BASH_VERSION:-}" ]; then
|
||||
if command -v bash >/dev/null 2>&1; then
|
||||
exec bash "$0" "$@"
|
||||
fi
|
||||
|
||||
if [ -f /etc/alpine-release ] && command -v apk >/dev/null 2>&1; then
|
||||
if [ "$(id -u)" -eq 0 ]; then
|
||||
apk add --no-cache bash
|
||||
elif command -v sudo >/dev/null 2>&1; then
|
||||
sudo apk add --no-cache bash
|
||||
elif command -v doas >/dev/null 2>&1; then
|
||||
doas apk add --no-cache bash
|
||||
else
|
||||
echo "❌ Alpine 安装需要 root 权限,或已配置 sudo/doas。" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
exec bash "$0" "$@"
|
||||
fi
|
||||
|
||||
echo "❌ 此安装脚本需要 Bash。" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# GitHub repo used for release downloads
|
||||
REPO="Sagit-chu/flux-panel"
|
||||
@@ -24,16 +51,51 @@ get_architecture() {
|
||||
|
||||
# 安装目录
|
||||
INSTALL_DIR="/etc/flux_agent"
|
||||
FLUX_AGENT_SYSTEMD_SERVICE_FILE="/etc/systemd/system/flux_agent.service"
|
||||
FLUX_AGENT_OPENRC_SERVICE_FILE="/etc/init.d/flux_agent"
|
||||
LEGACY_GOST_BINARY="/usr/local/bin/gost"
|
||||
LEGACY_GOST_CONFIG_DIR="/etc/gost"
|
||||
LEGACY_GOST_SERVICE_FILE_ETC="/etc/systemd/system/gost.service"
|
||||
LEGACY_GOST_SERVICE_FILE_LIB="/lib/systemd/system/gost.service"
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB="/usr/lib/systemd/system/gost.service"
|
||||
SERVICE_MANAGER="${SERVICE_MANAGER:-}"
|
||||
|
||||
# 镜像加速配置(可由面板传入或交互式询问)
|
||||
PROXY_ENABLED="${PROXY_ENABLED:-}"
|
||||
PROXY_URL="${PROXY_URL:-}"
|
||||
|
||||
ensure_alpine_runtime_dependencies() {
|
||||
[[ -f /etc/alpine-release ]] || return 0
|
||||
|
||||
local missing_packages=()
|
||||
local privileged_command=""
|
||||
|
||||
command -v curl >/dev/null 2>&1 || missing_packages+=(curl)
|
||||
[[ -f /etc/ssl/certs/ca-certificates.crt ]] || missing_packages+=(ca-certificates)
|
||||
|
||||
if [[ ${#missing_packages[@]} -eq 0 ]]; then
|
||||
return 0
|
||||
fi
|
||||
|
||||
if [[ $EUID -ne 0 ]]; then
|
||||
if command -v sudo >/dev/null 2>&1; then
|
||||
privileged_command="sudo"
|
||||
elif command -v doas >/dev/null 2>&1; then
|
||||
privileged_command="doas"
|
||||
else
|
||||
echo "❌ Alpine 安装需要 root 权限,或已配置 sudo/doas 来安装依赖: ${missing_packages[*]}。" >&2
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "📦 Alpine 缺少运行依赖,正在安装: ${missing_packages[*]}"
|
||||
if [[ -n "$privileged_command" ]]; then
|
||||
"$privileged_command" apk add --no-cache "${missing_packages[@]}"
|
||||
else
|
||||
apk add --no-cache "${missing_packages[@]}"
|
||||
fi
|
||||
}
|
||||
|
||||
# 镜像加速
|
||||
maybe_proxy_url() {
|
||||
local url="$1"
|
||||
@@ -139,6 +201,8 @@ build_download_url() {
|
||||
}
|
||||
|
||||
ensure_download_url_initialized() {
|
||||
ensure_alpine_runtime_dependencies || return 1
|
||||
|
||||
if [[ -n "${DOWNLOAD_URL:-}" ]]; then
|
||||
return 0
|
||||
fi
|
||||
@@ -256,6 +320,205 @@ write_flux_agent_config() {
|
||||
"$(json_escape "$SECRET")" > "$path"
|
||||
}
|
||||
|
||||
ensure_service_manager() {
|
||||
if [[ -n "$SERVICE_MANAGER" ]]; then
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd|openrc)
|
||||
return 0
|
||||
;;
|
||||
*)
|
||||
echo "❌ 不支持的服务管理器: $SERVICE_MANAGER" >&2
|
||||
return 1
|
||||
;;
|
||||
esac
|
||||
fi
|
||||
|
||||
# Alpine uses OpenRC even if a systemctl compatibility command happens to be installed.
|
||||
if [[ -f /etc/alpine-release ]]; then
|
||||
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
SERVICE_MANAGER="openrc"
|
||||
return 0
|
||||
fi
|
||||
elif command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
|
||||
SERVICE_MANAGER="systemd"
|
||||
return 0
|
||||
elif command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
|
||||
SERVICE_MANAGER="openrc"
|
||||
return 0
|
||||
fi
|
||||
|
||||
echo "❌ 未检测到受支持的服务管理器(systemd 或 OpenRC)。" >&2
|
||||
return 1
|
||||
}
|
||||
|
||||
flux_agent_service_exists() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
[[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" ]] || \
|
||||
systemctl list-units --full -all 2>/dev/null | grep -Fq "flux_agent.service"
|
||||
;;
|
||||
openrc)
|
||||
[[ -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
stop_flux_agent_service() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
systemctl stop flux_agent 2>/dev/null || true
|
||||
;;
|
||||
openrc)
|
||||
rc-service flux_agent stop 2>/dev/null || true
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
disable_flux_agent_service() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
systemctl disable flux_agent 2>/dev/null || true
|
||||
;;
|
||||
openrc)
|
||||
rc-update del flux_agent default 2>/dev/null || true
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
write_flux_agent_service() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
mkdir -p "$(dirname "$FLUX_AGENT_SYSTEMD_SERVICE_FILE")"
|
||||
cat > "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" <<EOF
|
||||
[Unit]
|
||||
Description=Flux_agent Proxy Service
|
||||
After=network.target
|
||||
|
||||
[Service]
|
||||
WorkingDirectory=$INSTALL_DIR
|
||||
ExecStart=$INSTALL_DIR/flux_agent
|
||||
Environment=GODEBUG=disablethp=1
|
||||
Restart=on-failure
|
||||
StandardOutput=null
|
||||
StandardError=null
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
EOF
|
||||
;;
|
||||
openrc)
|
||||
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
|
||||
cat > "$FLUX_AGENT_OPENRC_SERVICE_FILE" <<EOF
|
||||
#!/sbin/openrc-run
|
||||
|
||||
name="flux_agent"
|
||||
description="Flux_agent Proxy Service"
|
||||
command="$INSTALL_DIR/flux_agent"
|
||||
directory="$INSTALL_DIR"
|
||||
command_background="yes"
|
||||
pidfile="/run/\${RC_SVCNAME}.pid"
|
||||
output_log="/dev/null"
|
||||
error_log="/dev/null"
|
||||
|
||||
depend() {
|
||||
need net
|
||||
}
|
||||
EOF
|
||||
chmod +x "$FLUX_AGENT_OPENRC_SERVICE_FILE"
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
enable_and_start_flux_agent_service() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
systemctl daemon-reload
|
||||
systemctl enable flux_agent
|
||||
systemctl start flux_agent
|
||||
;;
|
||||
openrc)
|
||||
rc-update add flux_agent default
|
||||
rc-service flux_agent start
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
start_flux_agent_service() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
systemctl start flux_agent
|
||||
;;
|
||||
openrc)
|
||||
rc-service flux_agent start
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
flux_agent_service_is_active() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
systemctl is-active --quiet flux_agent
|
||||
;;
|
||||
openrc)
|
||||
rc-service flux_agent status >/dev/null 2>&1
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
flux_agent_service_status() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
systemctl is-active flux_agent 2>/dev/null || true
|
||||
;;
|
||||
openrc)
|
||||
rc-service flux_agent status 2>/dev/null || true
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
remove_flux_agent_service() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
rm -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE"
|
||||
systemctl daemon-reload 2>/dev/null || true
|
||||
;;
|
||||
openrc)
|
||||
rm -f "$FLUX_AGENT_OPENRC_SERVICE_FILE"
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
flux_agent_service_status_hint() {
|
||||
ensure_service_manager || return 1
|
||||
|
||||
case "$SERVICE_MANAGER" in
|
||||
systemd)
|
||||
echo "systemctl status flux_agent --no-pager"
|
||||
;;
|
||||
openrc)
|
||||
echo "rc-service flux_agent status"
|
||||
;;
|
||||
esac
|
||||
}
|
||||
|
||||
cleanup_legacy_gost_installation() {
|
||||
local matched_service_files=()
|
||||
local service_file=""
|
||||
@@ -278,7 +541,8 @@ cleanup_legacy_gost_installation() {
|
||||
return 0
|
||||
fi
|
||||
|
||||
if systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
if [[ "$SERVICE_MANAGER" == "systemd" ]] && \
|
||||
systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
|
||||
systemctl stop gost 2>/dev/null || true
|
||||
systemctl disable gost 2>/dev/null || true
|
||||
fi
|
||||
@@ -297,7 +561,7 @@ cleanup_legacy_gost_installation() {
|
||||
rm -f "$LEGACY_GOST_CONFIG_DIR/gost"
|
||||
fi
|
||||
|
||||
if [[ "$removed_service_file" == "1" ]]; then
|
||||
if [[ "$removed_service_file" == "1" && "$SERVICE_MANAGER" == "systemd" ]]; then
|
||||
systemctl daemon-reload 2>/dev/null || true
|
||||
fi
|
||||
}
|
||||
@@ -341,19 +605,21 @@ install_flux_agent() {
|
||||
|
||||
get_config_params
|
||||
|
||||
# 检查并安装 tcpkill
|
||||
# 检查并安装 tcpkill
|
||||
check_and_install_tcpkill
|
||||
|
||||
ensure_service_manager || exit 1
|
||||
|
||||
mkdir -p "$INSTALL_DIR"
|
||||
|
||||
local tmp_binary="$INSTALL_DIR/flux_agent.new"
|
||||
|
||||
# 停止并禁用已有服务
|
||||
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
|
||||
if flux_agent_service_exists; then
|
||||
echo "🔍 检测到已存在的flux_agent服务"
|
||||
systemctl stop flux_agent 2>/dev/null && echo "🛑 停止服务"
|
||||
systemctl disable flux_agent 2>/dev/null && echo "🚫 禁用自启"
|
||||
stop_flux_agent_service
|
||||
echo "🛑 停止服务"
|
||||
disable_flux_agent_service
|
||||
echo "🚫 禁用自启"
|
||||
fi
|
||||
|
||||
# 下载 flux_agent
|
||||
@@ -392,38 +658,21 @@ EOF
|
||||
# 加强权限
|
||||
chmod 600 "$INSTALL_DIR"/*.json
|
||||
|
||||
# 创建 systemd 服务
|
||||
SERVICE_FILE="/etc/systemd/system/flux_agent.service"
|
||||
cat > "$SERVICE_FILE" <<EOF
|
||||
[Unit]
|
||||
Description=Flux_agent Proxy Service
|
||||
After=network.target
|
||||
|
||||
[Service]
|
||||
WorkingDirectory=$INSTALL_DIR
|
||||
ExecStart=$INSTALL_DIR/flux_agent
|
||||
Restart=on-failure
|
||||
StandardOutput=null
|
||||
StandardError=null
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
EOF
|
||||
# 创建 systemd 或 OpenRC 服务
|
||||
write_flux_agent_service
|
||||
|
||||
# 启动服务
|
||||
systemctl daemon-reload
|
||||
systemctl enable flux_agent
|
||||
systemctl start flux_agent
|
||||
enable_and_start_flux_agent_service
|
||||
|
||||
# 检查状态
|
||||
echo "🔄 检查服务状态..."
|
||||
if systemctl is-active --quiet flux_agent; then
|
||||
if flux_agent_service_is_active; then
|
||||
echo "✅ 安装完成,flux_agent服务已启动并设置为开机启动。"
|
||||
echo "📁 配置目录: $INSTALL_DIR"
|
||||
echo "🔧 服务状态: $(systemctl is-active flux_agent)"
|
||||
echo "🔧 服务状态: $(flux_agent_service_status)"
|
||||
else
|
||||
echo "❌ flux_agent服务启动失败,请执行以下命令查看状态:"
|
||||
echo "systemctl status flux_agent --no-pager"
|
||||
flux_agent_service_status_hint
|
||||
fi
|
||||
}
|
||||
|
||||
@@ -443,6 +692,7 @@ update_flux_agent() {
|
||||
|
||||
# 检查并安装 tcpkill
|
||||
check_and_install_tcpkill
|
||||
ensure_service_manager || return 1
|
||||
|
||||
# 先下载新版本
|
||||
echo "⬇️ 下载最新版本..."
|
||||
@@ -455,9 +705,9 @@ update_flux_agent() {
|
||||
cleanup_legacy_gost_installation
|
||||
|
||||
# 停止服务
|
||||
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
|
||||
if flux_agent_service_exists; then
|
||||
echo "🛑 停止 flux_agent 服务..."
|
||||
systemctl stop flux_agent
|
||||
stop_flux_agent_service
|
||||
fi
|
||||
|
||||
# 替换文件
|
||||
@@ -469,7 +719,7 @@ update_flux_agent() {
|
||||
|
||||
# 重启服务
|
||||
echo "🔄 重启服务..."
|
||||
systemctl start flux_agent
|
||||
start_flux_agent_service
|
||||
|
||||
echo "✅ 更新完成,服务已重新启动。"
|
||||
}
|
||||
@@ -477,6 +727,7 @@ update_flux_agent() {
|
||||
# 卸载功能
|
||||
uninstall_flux_agent() {
|
||||
echo "🗑️ 开始卸载 flux_agent..."
|
||||
ensure_service_manager || return 1
|
||||
|
||||
read -p "确认卸载 flux_agent 吗?此操作将删除所有相关文件 (y/N): " confirm
|
||||
if [[ "$confirm" != "y" && "$confirm" != "Y" ]]; then
|
||||
@@ -485,15 +736,15 @@ uninstall_flux_agent() {
|
||||
fi
|
||||
|
||||
# 停止并禁用服务
|
||||
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
|
||||
if flux_agent_service_exists; then
|
||||
echo "🛑 停止并禁用服务..."
|
||||
systemctl stop flux_agent 2>/dev/null
|
||||
systemctl disable flux_agent 2>/dev/null
|
||||
stop_flux_agent_service
|
||||
disable_flux_agent_service
|
||||
fi
|
||||
|
||||
# 删除服务文件
|
||||
if [[ -f "/etc/systemd/system/flux_agent.service" ]]; then
|
||||
rm -f "/etc/systemd/system/flux_agent.service"
|
||||
if [[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]; then
|
||||
remove_flux_agent_service
|
||||
echo "🧹 删除服务文件"
|
||||
fi
|
||||
|
||||
@@ -503,9 +754,6 @@ uninstall_flux_agent() {
|
||||
echo "🧹 删除安装目录: $INSTALL_DIR"
|
||||
fi
|
||||
|
||||
# 重载 systemd
|
||||
systemctl daemon-reload
|
||||
|
||||
echo "✅ 卸载完成"
|
||||
}
|
||||
|
||||
|
||||
+319
-48
@@ -16,6 +16,9 @@ PINNED_VERSION=""
|
||||
# 镜像加速配置(可由面板传入或交互式询问)
|
||||
PROXY_ENABLED="${PROXY_ENABLED:-}"
|
||||
PROXY_URL="${PROXY_URL:-}"
|
||||
DEFAULT_PANEL_BACKEND_CONTAINER="flux-panel-backend"
|
||||
DEFAULT_PANEL_POSTGRES_CONTAINER="flux-panel-postgres"
|
||||
PANEL_SCRIPT_PATH="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)/$(basename "${BASH_SOURCE[0]}")"
|
||||
|
||||
# 镜像加速
|
||||
maybe_proxy_url() {
|
||||
@@ -304,8 +307,8 @@ get_env_var() {
|
||||
get_current_db_type() {
|
||||
local db_type database_url
|
||||
|
||||
db_type=$(get_env_var "DB_TYPE")
|
||||
database_url=$(get_env_var "DATABASE_URL")
|
||||
db_type=$(get_env_var "DB_TYPE" || true)
|
||||
database_url=$(get_env_var "DATABASE_URL" || true)
|
||||
|
||||
if [[ "$db_type" == "sqlite" ]]; then
|
||||
echo "sqlite"
|
||||
@@ -316,13 +319,289 @@ get_current_db_type() {
|
||||
fi
|
||||
}
|
||||
|
||||
get_container_compose_label() {
|
||||
local container="$1"
|
||||
local label="$2"
|
||||
local value
|
||||
|
||||
value=$(docker inspect -f "{{ index .Config.Labels \"$label\" }}" "$container" 2>/dev/null || true)
|
||||
if [[ "$value" == "<no value>" ]]; then
|
||||
value=""
|
||||
fi
|
||||
printf '%s' "$value"
|
||||
}
|
||||
|
||||
resolve_panel_deployment() {
|
||||
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
|
||||
local requested_dir="${PANEL_DEPLOY_DIR:-}"
|
||||
local label_dir label_project deploy_dir project_name
|
||||
|
||||
label_dir=$(get_container_compose_label "$backend_container" "com.docker.compose.project.working_dir")
|
||||
label_project=$(get_container_compose_label "$backend_container" "com.docker.compose.project")
|
||||
|
||||
if [[ -n "$requested_dir" ]]; then
|
||||
deploy_dir="$requested_dir"
|
||||
elif [[ -n "$label_dir" ]]; then
|
||||
deploy_dir="$label_dir"
|
||||
elif [[ -f ".env" && -f "docker-compose.yml" ]]; then
|
||||
deploy_dir=$(pwd -P)
|
||||
else
|
||||
echo "❌ 无法识别面板部署目录:未找到容器 Compose 标签,当前目录也没有完整部署配置"
|
||||
return 1
|
||||
fi
|
||||
|
||||
if [[ ! -d "$deploy_dir" ]]; then
|
||||
echo "❌ 面板部署目录不存在:$deploy_dir"
|
||||
return 1
|
||||
fi
|
||||
deploy_dir=$(cd "$deploy_dir" && pwd -P)
|
||||
if [[ -n "$label_dir" && -d "$label_dir" ]]; then
|
||||
label_dir=$(cd "$label_dir" && pwd -P)
|
||||
fi
|
||||
if [[ -n "$requested_dir" && -n "$label_dir" && "$deploy_dir" != "$label_dir" ]]; then
|
||||
echo "❌ 指定部署目录与运行中容器标签不一致:$deploy_dir != $label_dir"
|
||||
return 1
|
||||
fi
|
||||
|
||||
project_name="${COMPOSE_PROJECT_NAME:-}"
|
||||
if [[ -z "$project_name" && -n "$label_project" ]]; then
|
||||
project_name="$label_project"
|
||||
fi
|
||||
if [[ -n "$project_name" && ! "$project_name" =~ ^[a-z0-9][a-z0-9_-]*$ ]]; then
|
||||
echo "❌ Compose 项目名不合法:$project_name"
|
||||
return 1
|
||||
fi
|
||||
|
||||
PANEL_DEPLOY_DIR="$deploy_dir"
|
||||
PANEL_BACKEND_CONTAINER="$backend_container"
|
||||
PANEL_POSTGRES_CONTAINER="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
|
||||
PANEL_COMPOSE_PROJECT="$project_name"
|
||||
export PANEL_DEPLOY_DIR PANEL_BACKEND_CONTAINER PANEL_POSTGRES_CONTAINER PANEL_COMPOSE_PROJECT
|
||||
if [[ -n "$project_name" ]]; then
|
||||
COMPOSE_PROJECT_NAME="$project_name"
|
||||
export COMPOSE_PROJECT_NAME
|
||||
fi
|
||||
|
||||
cd "$PANEL_DEPLOY_DIR"
|
||||
echo "📁 部署目录:$PANEL_DEPLOY_DIR"
|
||||
if [[ -n "$PANEL_COMPOSE_PROJECT" ]]; then
|
||||
echo "📦 Compose 项目:$PANEL_COMPOSE_PROJECT"
|
||||
fi
|
||||
}
|
||||
|
||||
run_panel_compose() {
|
||||
$DOCKER_CMD "$@"
|
||||
}
|
||||
|
||||
validate_panel_update_environment() {
|
||||
local key value backend_port frontend_port configured_db_type db_type database_url postgres_password
|
||||
|
||||
if [[ ! -f ".env" ]]; then
|
||||
echo "❌ 部署目录缺少 .env,更新终止"
|
||||
return 1
|
||||
fi
|
||||
if [[ ! -f "docker-compose.yml" ]]; then
|
||||
echo "❌ 部署目录缺少 docker-compose.yml,更新终止"
|
||||
return 1
|
||||
fi
|
||||
|
||||
for key in JWT_SECRET BACKEND_PORT FRONTEND_PORT; do
|
||||
value=$(get_env_var "$key" || true)
|
||||
if [[ -z "${value//[[:space:]]/}" ]]; then
|
||||
echo "❌ .env 中 $key 缺失或为空,更新终止"
|
||||
return 1
|
||||
fi
|
||||
done
|
||||
|
||||
backend_port=$(get_env_var "BACKEND_PORT" || true)
|
||||
frontend_port=$(get_env_var "FRONTEND_PORT" || true)
|
||||
if [[ ! "$backend_port" =~ ^[0-9]+$ || "$backend_port" -lt 1 || "$backend_port" -gt 65535 ]]; then
|
||||
echo "❌ BACKEND_PORT 不是有效端口:$backend_port"
|
||||
return 1
|
||||
fi
|
||||
if [[ ! "$frontend_port" =~ ^[0-9]+$ || "$frontend_port" -lt 1 || "$frontend_port" -gt 65535 ]]; then
|
||||
echo "❌ FRONTEND_PORT 不是有效端口:$frontend_port"
|
||||
return 1
|
||||
fi
|
||||
|
||||
configured_db_type=$(get_env_var "DB_TYPE" || true)
|
||||
if [[ -n "$configured_db_type" && "$configured_db_type" != "sqlite" && "$configured_db_type" != "postgres" ]]; then
|
||||
echo "❌ DB_TYPE 仅支持 sqlite 或 postgres:$configured_db_type"
|
||||
return 1
|
||||
fi
|
||||
db_type=$(get_current_db_type)
|
||||
if [[ "$db_type" == "postgres" ]]; then
|
||||
database_url=$(get_env_var "DATABASE_URL" || true)
|
||||
postgres_password=$(get_env_var "POSTGRES_PASSWORD" || true)
|
||||
if [[ -z "${database_url//[[:space:]]/}" || -z "${postgres_password//[[:space:]]/}" ]]; then
|
||||
echo "❌ PostgreSQL 模式要求 DATABASE_URL 和 POSTGRES_PASSWORD 均非空"
|
||||
return 1
|
||||
fi
|
||||
fi
|
||||
|
||||
if ! run_panel_compose -f docker-compose.yml config -q; then
|
||||
echo "❌ 当前 Compose 配置校验失败,更新终止"
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
download_panel_update_compose() {
|
||||
local compose_url="$1"
|
||||
|
||||
UPDATE_COMPOSE_CANDIDATE=$(mktemp "$PANEL_DEPLOY_DIR/.docker-compose.yml.update.XXXXXX")
|
||||
if ! curl -fL -o "$UPDATE_COMPOSE_CANDIDATE" "$compose_url"; then
|
||||
rm -f "$UPDATE_COMPOSE_CANDIDATE"
|
||||
UPDATE_COMPOSE_CANDIDATE=""
|
||||
echo "❌ 下载最新 Compose 配置失败,更新终止"
|
||||
return 1
|
||||
fi
|
||||
if [[ ! -s "$UPDATE_COMPOSE_CANDIDATE" ]]; then
|
||||
rm -f "$UPDATE_COMPOSE_CANDIDATE"
|
||||
UPDATE_COMPOSE_CANDIDATE=""
|
||||
echo "❌ 下载的 Compose 配置为空,更新终止"
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
validate_panel_update_compose() {
|
||||
if ! FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" config -q; then
|
||||
echo "❌ 新 Compose 配置校验失败,更新终止"
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
backup_sqlite_for_update() (
|
||||
local destination="$1"
|
||||
local running paused="false"
|
||||
|
||||
cleanup_paused_backend() {
|
||||
if [[ "$paused" == "true" ]]; then
|
||||
docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null 2>&1 || true
|
||||
fi
|
||||
}
|
||||
trap cleanup_paused_backend EXIT
|
||||
|
||||
running=$(docker inspect -f '{{.State.Running}}' "$PANEL_BACKEND_CONTAINER" 2>/dev/null || true)
|
||||
if [[ "$running" == "true" ]]; then
|
||||
if ! docker pause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
|
||||
echo "❌ 暂停后端以创建一致性 SQLite 备份失败"
|
||||
return 1
|
||||
fi
|
||||
paused="true"
|
||||
fi
|
||||
|
||||
mkdir -p "$destination"
|
||||
if ! docker cp "$PANEL_BACKEND_CONTAINER:/app/data/." "$destination"; then
|
||||
echo "❌ SQLite 数据备份失败"
|
||||
return 1
|
||||
fi
|
||||
|
||||
if [[ "$paused" == "true" ]] && ! docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
|
||||
echo "❌ SQLite 备份完成,但恢复后端运行失败"
|
||||
return 1
|
||||
fi
|
||||
paused="false"
|
||||
)
|
||||
|
||||
backup_postgres_for_update() {
|
||||
local destination="$1"
|
||||
local postgres_db postgres_user
|
||||
|
||||
postgres_db=$(get_env_var "POSTGRES_DB" || true)
|
||||
postgres_user=$(get_env_var "POSTGRES_USER" || true)
|
||||
postgres_db=${postgres_db:-flux_panel}
|
||||
postgres_user=${postgres_user:-flux_panel}
|
||||
|
||||
if ! docker exec "$PANEL_POSTGRES_CONTAINER" pg_dump -U "$postgres_user" "$postgres_db" > "$destination"; then
|
||||
rm -f "$destination"
|
||||
echo "❌ PostgreSQL 数据备份失败"
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
create_panel_update_backup() {
|
||||
local db_type="$1"
|
||||
local timestamp
|
||||
|
||||
timestamp=$(date '+%Y%m%d-%H%M%S')
|
||||
if ! (umask 077 && mkdir -p "$PANEL_DEPLOY_DIR/backups"); then
|
||||
echo "❌ 创建更新备份目录失败"
|
||||
return 1
|
||||
fi
|
||||
UPDATE_BACKUP_DIR=$(umask 077 && mktemp -d "$PANEL_DEPLOY_DIR/backups/panel-update-$timestamp.XXXXXX") || {
|
||||
echo "❌ 创建更新备份目录失败"
|
||||
return 1
|
||||
}
|
||||
if ! cp -p .env docker-compose.yml "$UPDATE_BACKUP_DIR/"; then
|
||||
echo "❌ 备份面板配置失败"
|
||||
return 1
|
||||
fi
|
||||
|
||||
if [[ "$db_type" == "postgres" ]]; then
|
||||
backup_postgres_for_update "$UPDATE_BACKUP_DIR/postgres.sql" || return 1
|
||||
else
|
||||
backup_sqlite_for_update "$UPDATE_BACKUP_DIR/sqlite" || return 1
|
||||
fi
|
||||
echo "💾 更新备份:$UPDATE_BACKUP_DIR"
|
||||
}
|
||||
|
||||
pull_panel_update_images() {
|
||||
local db_type="$1"
|
||||
|
||||
if [[ "$db_type" == "postgres" ]]; then
|
||||
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend postgres
|
||||
else
|
||||
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend
|
||||
fi
|
||||
}
|
||||
|
||||
activate_panel_update_files() {
|
||||
if ! chmod 0644 "$UPDATE_COMPOSE_CANDIDATE" || ! mv "$UPDATE_COMPOSE_CANDIDATE" docker-compose.yml; then
|
||||
echo "❌ 替换 Compose 配置失败"
|
||||
return 1
|
||||
fi
|
||||
UPDATE_COMPOSE_CANDIDATE=""
|
||||
if ! upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"; then
|
||||
echo "❌ 更新版本配置失败"
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
start_panel_after_update() {
|
||||
local db_type="$1"
|
||||
|
||||
if [[ "$db_type" == "postgres" ]]; then
|
||||
run_panel_compose up -d postgres || return 1
|
||||
wait_for_postgres_healthy || return 1
|
||||
fi
|
||||
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
|
||||
wait_for_backend_healthy
|
||||
}
|
||||
|
||||
rollback_panel_update() {
|
||||
local db_type="$1"
|
||||
|
||||
echo "↩️ 正在恢复更新前配置..."
|
||||
if ! cp -p "$UPDATE_BACKUP_DIR/.env" .env || ! cp -p "$UPDATE_BACKUP_DIR/docker-compose.yml" docker-compose.yml; then
|
||||
echo "❌ 配置回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
|
||||
return 1
|
||||
fi
|
||||
if [[ "$db_type" == "postgres" ]]; then
|
||||
run_panel_compose up -d postgres || return 1
|
||||
wait_for_postgres_healthy || return 1
|
||||
fi
|
||||
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
|
||||
wait_for_backend_healthy
|
||||
}
|
||||
|
||||
wait_for_postgres_healthy() {
|
||||
local pg_health
|
||||
local postgres_container="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
|
||||
|
||||
echo "🔍 检查 PostgreSQL 服务状态..."
|
||||
for i in {1..90}; do
|
||||
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-postgres$"; then
|
||||
pg_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo "unknown")
|
||||
if docker ps --format "{{.Names}}" | grep -Fxq "$postgres_container"; then
|
||||
pg_health=$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo "unknown")
|
||||
if [[ "$pg_health" == "healthy" ]]; then
|
||||
echo "✅ PostgreSQL 服务健康检查通过"
|
||||
return 0
|
||||
@@ -335,7 +614,7 @@ wait_for_postgres_healthy() {
|
||||
|
||||
if [ $i -eq 90 ]; then
|
||||
echo "❌ PostgreSQL 启动超时(90秒)"
|
||||
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo '容器不存在')"
|
||||
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo '容器不存在')"
|
||||
return 1
|
||||
fi
|
||||
|
||||
@@ -348,11 +627,12 @@ wait_for_postgres_healthy() {
|
||||
|
||||
wait_for_backend_healthy() {
|
||||
local backend_health
|
||||
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
|
||||
|
||||
echo "🔍 检查后端服务状态..."
|
||||
for i in {1..90}; do
|
||||
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-backend$"; then
|
||||
backend_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo "unknown")
|
||||
if docker ps --format "{{.Names}}" | grep -Fxq "$backend_container"; then
|
||||
backend_health=$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo "unknown")
|
||||
if [[ "$backend_health" == "healthy" ]]; then
|
||||
echo "✅ 后端服务健康检查通过"
|
||||
return 0
|
||||
@@ -365,7 +645,7 @@ wait_for_backend_healthy() {
|
||||
|
||||
if [ $i -eq 90 ]; then
|
||||
echo "❌ 后端服务启动超时(90秒)"
|
||||
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo '容器不存在')"
|
||||
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo '容器不存在')"
|
||||
return 1
|
||||
fi
|
||||
|
||||
@@ -380,9 +660,8 @@ wait_for_backend_healthy() {
|
||||
delete_self() {
|
||||
echo ""
|
||||
echo "🗑️ 操作已完成,正在清理脚本文件..."
|
||||
SCRIPT_PATH="$(readlink -f "$0" 2>/dev/null || realpath "$0" 2>/dev/null || echo "$0")"
|
||||
sleep 1
|
||||
rm -f "$SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
|
||||
rm -f "$PANEL_SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
|
||||
}
|
||||
|
||||
|
||||
@@ -486,12 +765,11 @@ EOF
|
||||
# 更新功能
|
||||
update_panel() {
|
||||
echo "🔄 开始更新面板..."
|
||||
ask_proxy_config
|
||||
check_docker
|
||||
resolve_panel_deployment || return 1
|
||||
ask_proxy_config
|
||||
validate_panel_update_environment || return 1
|
||||
|
||||
if [[ ! -f ".env" ]]; then
|
||||
echo "⚠️ 未找到 .env,默认按 SQLite 模式更新"
|
||||
fi
|
||||
CURRENT_DB_TYPE=$(get_current_db_type)
|
||||
echo "🗄️ 当前数据库类型:$CURRENT_DB_TYPE"
|
||||
|
||||
@@ -502,52 +780,45 @@ update_panel() {
|
||||
}
|
||||
echo "🆕 最新版本:$LATEST_VERSION"
|
||||
set_compose_urls_by_version "$LATEST_VERSION"
|
||||
upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"
|
||||
|
||||
echo "🔽 下载最新配置文件..."
|
||||
DOCKER_COMPOSE_URL=$(get_docker_compose_url)
|
||||
echo "📡 选择配置文件:$(basename "$DOCKER_COMPOSE_URL")"
|
||||
curl -L -o docker-compose.yml "$DOCKER_COMPOSE_URL"
|
||||
echo "✅ 下载完成"
|
||||
|
||||
# 自动检测并配置 IPv6 支持
|
||||
if check_ipv6_support; then
|
||||
echo "🚀 系统支持 IPv6,自动启用 IPv6 配置..."
|
||||
configure_docker_ipv6
|
||||
download_panel_update_compose "$DOCKER_COMPOSE_URL" || return 1
|
||||
if ! validate_panel_update_compose; then
|
||||
rm -f "$UPDATE_COMPOSE_CANDIDATE"
|
||||
UPDATE_COMPOSE_CANDIDATE=""
|
||||
return 1
|
||||
fi
|
||||
|
||||
# 先发送 SIGTERM 信号,让应用优雅关闭
|
||||
docker stop -t 30 flux-panel-backend 2>/dev/null || true
|
||||
docker stop -t 10 vite-frontend 2>/dev/null || true
|
||||
|
||||
# 等待 WAL 文件同步
|
||||
echo "⏳ 等待数据同步..."
|
||||
sleep 5
|
||||
|
||||
# 然后再完全停止
|
||||
$DOCKER_CMD down
|
||||
echo "💾 备份当前配置和数据库..."
|
||||
if ! create_panel_update_backup "$CURRENT_DB_TYPE"; then
|
||||
rm -f "$UPDATE_COMPOSE_CANDIDATE"
|
||||
UPDATE_COMPOSE_CANDIDATE=""
|
||||
return 1
|
||||
fi
|
||||
|
||||
echo "⬇️ 拉取最新镜像..."
|
||||
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
|
||||
$DOCKER_CMD pull backend frontend postgres
|
||||
else
|
||||
$DOCKER_CMD pull backend frontend
|
||||
if ! pull_panel_update_images "$CURRENT_DB_TYPE"; then
|
||||
rm -f "$UPDATE_COMPOSE_CANDIDATE"
|
||||
UPDATE_COMPOSE_CANDIDATE=""
|
||||
echo "❌ 镜像拉取失败,现有服务保持运行"
|
||||
return 1
|
||||
fi
|
||||
|
||||
echo "🚀 启动更新后的服务..."
|
||||
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
|
||||
$DOCKER_CMD up -d postgres
|
||||
wait_for_postgres_healthy
|
||||
$DOCKER_CMD up -d backend frontend
|
||||
else
|
||||
$DOCKER_CMD up -d backend frontend
|
||||
if ! activate_panel_update_files; then
|
||||
rollback_panel_update "$CURRENT_DB_TYPE" || true
|
||||
return 1
|
||||
fi
|
||||
|
||||
# 等待服务启动
|
||||
echo "⏳ 等待服务启动..."
|
||||
|
||||
if ! wait_for_backend_healthy; then
|
||||
echo "🛑 更新终止"
|
||||
echo "🚀 原地重建更新后的服务..."
|
||||
if ! start_panel_after_update "$CURRENT_DB_TYPE"; then
|
||||
echo "🛑 新版本启动失败"
|
||||
if rollback_panel_update "$CURRENT_DB_TYPE"; then
|
||||
echo "✅ 已恢复更新前版本"
|
||||
else
|
||||
echo "❌ 自动回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
|
||||
fi
|
||||
return 1
|
||||
fi
|
||||
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
# Run only in a disposable Alpine container: this exercises real OpenRC services.
|
||||
# docker run --rm -v "$PWD:/workspace:ro" alpine:3.22 sh /workspace/test-install-scripts-alpine.sh
|
||||
[ -f /etc/alpine-release ] && [ "$(id -u)" = 0 ] || {
|
||||
echo "Run this test as root in a disposable Alpine container." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
ROOT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
|
||||
INSTALL_SCRIPT=${1:-"$ROOT_DIR/install.sh"}
|
||||
TEST_DIR=$(mktemp -d)
|
||||
export TEST_DIR
|
||||
|
||||
cleanup() {
|
||||
rc-service flux_agent stop >/dev/null 2>&1 || true
|
||||
rc-update del flux_agent default >/dev/null 2>&1 || true
|
||||
rm -f /etc/init.d/flux_agent
|
||||
rm -rf "$TEST_DIR"
|
||||
}
|
||||
trap cleanup EXIT
|
||||
|
||||
# Docker supplies the network; an empty interfaces file lets OpenRC register
|
||||
# that dependency without changing the container's network configuration.
|
||||
apk add --no-cache openrc
|
||||
mkdir -p /run/openrc /etc/network
|
||||
touch /run/openrc/softlevel /etc/network/interfaces
|
||||
rc-service networking start
|
||||
|
||||
# Fail immediately if any path accidentally calls systemd on Alpine.
|
||||
mkdir -p "$TEST_DIR/bin"
|
||||
cat > "$TEST_DIR/bin/systemctl" <<'EOF'
|
||||
#!/bin/sh
|
||||
touch "$TEST_DIR/systemctl-called"
|
||||
exit 1
|
||||
EOF
|
||||
chmod +x "$TEST_DIR/bin/systemctl"
|
||||
export PATH="$TEST_DIR/bin:$PATH"
|
||||
|
||||
cat > "$TEST_DIR/agent" <<'EOF'
|
||||
#!/bin/sh
|
||||
if [ "${1:-}" = -V ]; then
|
||||
echo "Alpine installer test agent"
|
||||
exit 0
|
||||
fi
|
||||
pwd > ../working-directory
|
||||
exec sleep 300
|
||||
EOF
|
||||
|
||||
# Keep the installer bootstrap, service detection and all lifecycle operations.
|
||||
# Substitute only the network download and optional tcpkill dependency.
|
||||
sed '/^# 执行主函数$/,$d' "$INSTALL_SCRIPT" > "$TEST_DIR/install.sh"
|
||||
cat >> "$TEST_DIR/install.sh" <<'EOF'
|
||||
INSTALL_DIR="$TEST_DIR/flux_agent"
|
||||
ensure_download_url_initialized() {
|
||||
ensure_alpine_runtime_dependencies || return 1
|
||||
DOWNLOAD_URL="file://$TEST_DIR/agent"
|
||||
}
|
||||
check_and_install_tcpkill() { :; }
|
||||
main
|
||||
EOF
|
||||
chmod +x "$TEST_DIR/install.sh"
|
||||
cp "$TEST_DIR/install.sh" "$TEST_DIR/manage.sh"
|
||||
PROXY_ENABLED=false "$TEST_DIR/install.sh" -a http://127.0.0.1:9 -s alpine-test
|
||||
|
||||
rc-service flux_agent status
|
||||
test -L /etc/runlevels/default/flux_agent
|
||||
test "$(cat "$TEST_DIR/working-directory")" = "$TEST_DIR/flux_agent"
|
||||
test -s "$TEST_DIR/flux_agent/config.json"
|
||||
rc-service flux_agent restart
|
||||
rc-service flux_agent status
|
||||
|
||||
# Update must retain config and restart through OpenRC.
|
||||
cp "$TEST_DIR/flux_agent/config.json" "$TEST_DIR/config-before-update.json"
|
||||
cp "$TEST_DIR/manage.sh" "$TEST_DIR/update.sh"
|
||||
printf '2\n' | PROXY_ENABLED=false "$TEST_DIR/update.sh"
|
||||
cmp "$TEST_DIR/config-before-update.json" "$TEST_DIR/flux_agent/config.json"
|
||||
rc-service flux_agent status
|
||||
|
||||
printf '3\ny\n' | "$TEST_DIR/manage.sh"
|
||||
test ! -e /etc/init.d/flux_agent
|
||||
test ! -e /etc/runlevels/default/flux_agent
|
||||
test ! -e "$TEST_DIR/flux_agent"
|
||||
test ! -e "$TEST_DIR/systemctl-called"
|
||||
echo "Alpine installer lifecycle tests passed"
|
||||
@@ -90,6 +90,7 @@ test_update_flux_agent_asks_for_proxy_config() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
SERVICE_MANAGER="systemd"
|
||||
INSTALL_DIR=$(mktemp -d)
|
||||
cat > "$INSTALL_DIR/flux_agent" <<'EOF'
|
||||
#!/bin/bash
|
||||
@@ -165,6 +166,7 @@ test_install_flux_agent_preserves_legacy_gost_when_download_fails() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
SERVICE_MANAGER="systemd"
|
||||
INSTALL_DIR=$(mktemp -d)
|
||||
cat > "$INSTALL_DIR/flux_agent" <<'EOF'
|
||||
#!/bin/bash
|
||||
@@ -199,6 +201,7 @@ test_update_flux_agent_preserves_legacy_gost_when_download_fails() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
SERVICE_MANAGER="systemd"
|
||||
INSTALL_DIR=$(mktemp -d)
|
||||
cat > "$INSTALL_DIR/flux_agent" <<'EOF'
|
||||
#!/bin/bash
|
||||
@@ -236,7 +239,9 @@ test_install_flux_agent_writes_json_safe_config() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
SERVICE_MANAGER="systemd"
|
||||
INSTALL_DIR=$(mktemp -d)
|
||||
FLUX_AGENT_SYSTEMD_SERVICE_FILE="$INSTALL_DIR/flux_agent.service"
|
||||
SERVER_ADDR='panel"addr'
|
||||
SECRET='sec\ret"1'
|
||||
DOWNLOAD_URL="https://example.com/gost"
|
||||
@@ -272,6 +277,119 @@ EOF
|
||||
local expected=$'{\n "addr": "panel\\"addr",\n "secret": "sec\\\\ret\\"1"\n}'
|
||||
|
||||
assert_equals "$expected" "$actual" "install_flux_agent should JSON-escape config values"
|
||||
grep -Fqx 'Environment=GODEBUG=disablethp=1' "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || \
|
||||
fail "systemd service should disable transparent huge pages for the Go heap"
|
||||
)
|
||||
|
||||
test_install_script_bootstraps_bash_for_alpine() (
|
||||
set -euo pipefail
|
||||
|
||||
local shebang
|
||||
shebang=$(head -n 1 "$ROOT_DIR/install.sh")
|
||||
assert_equals "#!/bin/sh" "$shebang" "install.sh should start with Alpine's default shell"
|
||||
|
||||
grep -Fq 'apk add --no-cache bash' "$ROOT_DIR/install.sh" || \
|
||||
fail "install.sh should bootstrap Bash through apk on Alpine"
|
||||
)
|
||||
|
||||
test_install_flux_agent_uses_openrc() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
local temp_root
|
||||
temp_root=$(mktemp -d)
|
||||
INSTALL_DIR="$temp_root/flux_agent"
|
||||
FLUX_AGENT_OPENRC_SERVICE_FILE="$temp_root/init.d/flux_agent"
|
||||
SERVICE_MANAGER="openrc"
|
||||
SERVER_ADDR="panel.example.com:443"
|
||||
SECRET="secret"
|
||||
DOWNLOAD_URL="https://example.com/gost"
|
||||
|
||||
local rc_service_calls=""
|
||||
local rc_update_calls=""
|
||||
|
||||
ask_proxy_config() { :; }
|
||||
ensure_download_url_initialized() { :; }
|
||||
get_config_params() { :; }
|
||||
check_and_install_tcpkill() { :; }
|
||||
cleanup_legacy_gost_installation() { :; }
|
||||
curl() {
|
||||
local output=""
|
||||
while [[ $# -gt 0 ]]; do
|
||||
if [[ "$1" == "-o" ]]; then
|
||||
output="$2"
|
||||
shift 2
|
||||
continue
|
||||
fi
|
||||
shift
|
||||
done
|
||||
|
||||
cat > "$output" <<'EOF'
|
||||
#!/bin/sh
|
||||
echo "new version"
|
||||
EOF
|
||||
chmod +x "$output"
|
||||
}
|
||||
rc-service() {
|
||||
rc_service_calls+=$'\n'"$*"
|
||||
if [[ "$2" == "status" ]]; then
|
||||
echo "status: started"
|
||||
fi
|
||||
return 0
|
||||
}
|
||||
rc-update() {
|
||||
rc_update_calls+=$'\n'"$*"
|
||||
return 0
|
||||
}
|
||||
|
||||
install_flux_agent >/dev/null
|
||||
|
||||
[[ -x "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]] || fail "OpenRC service file should be executable"
|
||||
grep -Fq '#!/sbin/openrc-run' "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
|
||||
fail "OpenRC service should use openrc-run"
|
||||
grep -Fq "command=\"$INSTALL_DIR/flux_agent\"" "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
|
||||
fail "OpenRC service should launch the installed flux_agent binary"
|
||||
grep -Fq 'command_background="yes"' "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
|
||||
fail "OpenRC service should run flux_agent in the background"
|
||||
if command -v openrc-run >/dev/null 2>&1; then
|
||||
"$FLUX_AGENT_OPENRC_SERVICE_FILE" describe >/dev/null 2>&1
|
||||
fi
|
||||
[[ "$rc_update_calls" == *"add flux_agent default"* ]] || \
|
||||
fail "OpenRC install should enable flux_agent in the default runlevel"
|
||||
[[ "$rc_service_calls" == *"start"* ]] || fail "OpenRC install should start flux_agent"
|
||||
[[ "$rc_service_calls" == *"status"* ]] || fail "OpenRC install should verify flux_agent status"
|
||||
)
|
||||
|
||||
test_remove_flux_agent_service_uses_openrc() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
local temp_root
|
||||
temp_root=$(mktemp -d)
|
||||
SERVICE_MANAGER="openrc"
|
||||
FLUX_AGENT_OPENRC_SERVICE_FILE="$temp_root/init.d/flux_agent"
|
||||
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
|
||||
: > "$FLUX_AGENT_OPENRC_SERVICE_FILE"
|
||||
|
||||
local rc_service_calls=""
|
||||
local rc_update_calls=""
|
||||
rc-service() {
|
||||
rc_service_calls+=$'\n'"$*"
|
||||
return 0
|
||||
}
|
||||
rc-update() {
|
||||
rc_update_calls+=$'\n'"$*"
|
||||
return 0
|
||||
}
|
||||
|
||||
stop_flux_agent_service
|
||||
disable_flux_agent_service
|
||||
remove_flux_agent_service
|
||||
|
||||
[[ "$rc_service_calls" == *"stop"* ]] || fail "OpenRC uninstall should stop flux_agent"
|
||||
[[ "$rc_update_calls" == *"del flux_agent default"* ]] || \
|
||||
fail "OpenRC uninstall should remove flux_agent from the default runlevel"
|
||||
[[ ! -e "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]] || fail "OpenRC uninstall should remove its service file"
|
||||
)
|
||||
|
||||
test_cleanup_legacy_gost_installation_removes_service_and_binary() (
|
||||
@@ -283,6 +401,7 @@ test_cleanup_legacy_gost_installation_removes_service_and_binary() (
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="systemd"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
|
||||
[Unit]
|
||||
Description=Gost Proxy Service
|
||||
@@ -326,6 +445,7 @@ test_cleanup_legacy_gost_installation_preserves_unrelated_gost() (
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="systemd"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<'EOF'
|
||||
[Unit]
|
||||
Description=Unrelated Gost Service
|
||||
@@ -353,6 +473,35 @@ EOF
|
||||
[[ "$systemctl_calls" != *"disable gost"* ]] || fail "cleanup_legacy_gost_installation should not disable unrelated gost services"
|
||||
)
|
||||
|
||||
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
|
||||
LEGACY_GOST_SERVICE_FILE_ETC=$(mktemp)
|
||||
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
|
||||
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
|
||||
SERVICE_MANAGER="openrc"
|
||||
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
|
||||
[Unit]
|
||||
WorkingDirectory=$LEGACY_GOST_CONFIG_DIR
|
||||
ExecStart=$LEGACY_GOST_CONFIG_DIR/gost
|
||||
EOF
|
||||
: > "$LEGACY_GOST_CONFIG_DIR/config.json"
|
||||
: > "$LEGACY_GOST_CONFIG_DIR/gost.json"
|
||||
|
||||
local systemctl_calls=""
|
||||
systemctl() {
|
||||
systemctl_calls+=$'\n'"$*"
|
||||
return 1
|
||||
}
|
||||
|
||||
cleanup_legacy_gost_installation >/dev/null
|
||||
|
||||
[[ -z "$systemctl_calls" ]] || fail "OpenRC cleanup should not invoke systemctl"
|
||||
[[ ! -e "$LEGACY_GOST_SERVICE_FILE_ETC" ]] || fail "OpenRC cleanup should still remove the legacy service file"
|
||||
)
|
||||
|
||||
test_install_script_accepts_proxy_url_env_without_prompt() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/install.sh"
|
||||
@@ -409,6 +558,7 @@ test_update_panel_asks_for_proxy_config() (
|
||||
load_script_without_main "$ROOT_DIR/panel_install.sh"
|
||||
|
||||
local ask_called="0"
|
||||
local calls=""
|
||||
|
||||
ask_proxy_config() {
|
||||
ask_called="1"
|
||||
@@ -421,6 +571,15 @@ test_update_panel_asks_for_proxy_config() (
|
||||
DOCKER_CMD="true"
|
||||
}
|
||||
|
||||
resolve_panel_deployment() {
|
||||
calls+=" resolve"
|
||||
PANEL_DEPLOY_DIR=$(mktemp -d)
|
||||
}
|
||||
|
||||
validate_panel_update_environment() {
|
||||
calls+=" validate-current"
|
||||
}
|
||||
|
||||
get_current_db_type() {
|
||||
echo "sqlite"
|
||||
}
|
||||
@@ -429,13 +588,33 @@ test_update_panel_asks_for_proxy_config() (
|
||||
echo "v-test"
|
||||
}
|
||||
|
||||
upsert_env_var() { :; }
|
||||
download_panel_update_compose() {
|
||||
calls+=" download"
|
||||
UPDATE_COMPOSE_CANDIDATE="candidate"
|
||||
}
|
||||
|
||||
validate_panel_update_compose() {
|
||||
calls+=" validate-new"
|
||||
}
|
||||
|
||||
create_panel_update_backup() {
|
||||
calls+=" backup"
|
||||
UPDATE_BACKUP_DIR="backup"
|
||||
}
|
||||
|
||||
pull_panel_update_images() {
|
||||
calls+=" pull"
|
||||
}
|
||||
|
||||
activate_panel_update_files() {
|
||||
calls+=" activate"
|
||||
}
|
||||
|
||||
start_panel_after_update() {
|
||||
calls+=" start"
|
||||
}
|
||||
|
||||
check_ipv6_support() { return 1; }
|
||||
configure_docker_ipv6() { :; }
|
||||
docker() { return 0; }
|
||||
wait_for_backend_healthy() { return 0; }
|
||||
sleep() { :; }
|
||||
curl() { :; }
|
||||
|
||||
update_panel >/dev/null
|
||||
|
||||
@@ -444,6 +623,75 @@ test_update_panel_asks_for_proxy_config() (
|
||||
"https://github.com/${REPO}/releases/download/v-test/docker-compose-v4.yml" \
|
||||
"$DOCKER_COMPOSEV4_URL" \
|
||||
"update_panel should honor the prompted proxy choice"
|
||||
assert_equals \
|
||||
" resolve validate-current download validate-new backup pull activate start" \
|
||||
"$calls" \
|
||||
"update_panel should validate, back up, and pull before activating the update"
|
||||
)
|
||||
|
||||
test_resolve_panel_deployment_uses_container_labels() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/panel_install.sh"
|
||||
|
||||
local expected_dir
|
||||
expected_dir=$(mktemp -d)
|
||||
expected_dir=$(cd "$expected_dir" && pwd -P)
|
||||
printf 'JWT_SECRET=test\nBACKEND_PORT=6365\nFRONTEND_PORT=6366\nDB_TYPE=sqlite\n' > "$expected_dir/.env"
|
||||
printf 'services: {}\n' > "$expected_dir/docker-compose.yml"
|
||||
|
||||
unset PANEL_DEPLOY_DIR COMPOSE_PROJECT_NAME
|
||||
docker() {
|
||||
if [[ "$1" == "inspect" && "$2" == "-f" ]]; then
|
||||
case "$3" in
|
||||
*working_dir*) printf '%s\n' "$expected_dir" ;;
|
||||
*) printf '%s\n' "flvx-panel" ;;
|
||||
esac
|
||||
fi
|
||||
}
|
||||
|
||||
resolve_panel_deployment >/dev/null
|
||||
|
||||
assert_equals "$expected_dir" "$PANEL_DEPLOY_DIR" "update should discover the Compose working directory from container labels"
|
||||
assert_equals "flvx-panel" "$COMPOSE_PROJECT_NAME" "update should preserve the existing Compose project name"
|
||||
assert_equals "$expected_dir" "$(pwd -P)" "update should run from the discovered deployment directory"
|
||||
)
|
||||
|
||||
test_validate_panel_update_environment_rejects_empty_required_values() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/panel_install.sh"
|
||||
|
||||
local deploy_dir rc="0"
|
||||
deploy_dir=$(mktemp -d)
|
||||
printf 'JWT_SECRET=\nBACKEND_PORT=6365\nFRONTEND_PORT=6366\nDB_TYPE=sqlite\n' > "$deploy_dir/.env"
|
||||
printf 'services: {}\n' > "$deploy_dir/docker-compose.yml"
|
||||
cd "$deploy_dir"
|
||||
DOCKER_CMD="true"
|
||||
|
||||
validate_panel_update_environment >/dev/null || rc="$?"
|
||||
|
||||
assert_equals "1" "$rc" "update should reject an empty JWT secret before touching the running service"
|
||||
)
|
||||
|
||||
test_backup_sqlite_for_update_pauses_copies_and_unpauses() (
|
||||
set -euo pipefail
|
||||
load_script_without_main "$ROOT_DIR/panel_install.sh"
|
||||
|
||||
local destination calls log_path
|
||||
destination=$(mktemp -d)/sqlite
|
||||
log_path=$(mktemp)
|
||||
PANEL_BACKEND_CONTAINER="flux-panel-backend"
|
||||
docker() {
|
||||
if [[ "$1" == "inspect" ]]; then
|
||||
printf 'true\n'
|
||||
return 0
|
||||
fi
|
||||
printf ' %s' "$1" >> "$log_path"
|
||||
}
|
||||
|
||||
backup_sqlite_for_update "$destination" >/dev/null
|
||||
calls=$(cat "$log_path")
|
||||
|
||||
assert_equals " pause cp unpause" "$calls" "SQLite backup should pause writes, copy the database volume, then resume the backend"
|
||||
)
|
||||
|
||||
test_panel_install_script_accepts_proxy_url_env_without_prompt() (
|
||||
@@ -504,14 +752,21 @@ test_update_flux_agent_skips_proxy_prompt_when_not_installed
|
||||
test_install_flux_agent_preserves_legacy_gost_when_download_fails
|
||||
test_update_flux_agent_preserves_legacy_gost_when_download_fails
|
||||
test_install_flux_agent_writes_json_safe_config
|
||||
test_install_script_bootstraps_bash_for_alpine
|
||||
test_install_flux_agent_uses_openrc
|
||||
test_remove_flux_agent_service_uses_openrc
|
||||
test_cleanup_legacy_gost_installation_removes_service_and_binary
|
||||
test_cleanup_legacy_gost_installation_preserves_unrelated_gost
|
||||
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc
|
||||
test_install_script_accepts_proxy_url_env_without_prompt
|
||||
test_panel_install_script_can_disable_proxy
|
||||
test_panel_install_script_recomputes_compose_urls_after_prompt
|
||||
test_update_panel_asks_for_proxy_config
|
||||
test_resolve_panel_deployment_uses_container_labels
|
||||
test_validate_panel_update_environment_rejects_empty_required_values
|
||||
test_backup_sqlite_for_update_pauses_copies_and_unpauses
|
||||
test_panel_install_script_uses_default_proxy
|
||||
test_panel_install_script_accepts_proxy_url_env_without_prompt
|
||||
test_panel_install_script_defaults_proxy_on_eof
|
||||
|
||||
echo "install script proxy tests passed"
|
||||
echo "install script tests passed"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user