Compare commits

...

25 Commits

Author SHA1 Message Date
sagitchu c05ff086e5 feat(monitor): show backup tunnel latencies 2026-08-03 10:13:00 +08:00
sagit ae370382d3 feat(agent): support Alpine installation (#534)
Add Alpine bootstrap and OpenRC lifecycle support to the agent installer.

Closes #527
2026-07-31 17:02:52 +08:00
sagit e112d81697 fix: bound and configure tunnel quality probes (#533)
Closes #528 and #532.
2026-07-31 15:37:26 +08:00
sagit 11a27d3c67 feat: add per-rule traffic reset (#526)
Reset only the selected forwarding rule's displayed upload/download usage without affecting user totals, tunnel quotas, historical statistics, nftables baselines, or running services.

Closes #523
2026-07-10 17:05:20 +08:00
sagit 8e513a1bae Fix nftables node updates after panel restart (#525)
## Summary
- Skip agent protocol SetProtocol commands when updating nftables nodes.
- Preserve existing nftables SSH credentials when edit forms omit secret
fields.
- Add regression coverage for nftables updates and SSH config
persistence across repository reopen.

## Test Plan
- rtk go test ./...
2026-07-02 10:28:00 +08:00
sagitchu f98be845d3 fix: skip agent protocol updates for nftables nodes 2026-07-02 10:26:10 +08:00
sagit 0c2acfdd8a Fix nftables recovery and diagnostics (#524)
## Summary
- Reconcile nftables nodes when background jobs start so rules are
restored after server reboot.
- Collect nftables traffic immediately at startup and every 30 seconds
by default.
- Return nftables rule binding status in forward diagnostics and cover
it with regression tests.

## Test Plan
- `cd go-backend && go test ./...`
- `cd go-backend && make build`
2026-06-30 17:11:10 +08:00
sagitchu 777db8767f fix nftables recovery and diagnostics 2026-06-30 17:07:50 +08:00
sagit 82f6047506 fix(forward): split proxy protocol directions
Closes #520
2026-06-21 21:03:47 +08:00
sagitchu 3ce320da5a fix(nftables): harden traffic accounting edges 2026-06-07 11:55:59 +08:00
sagitchu 7ab0db29ae fix(nftables): clean counter state on forward delete 2026-06-07 11:55:59 +08:00
sagitchu 006ea97200 feat(nftables): ingest traffic counters 2026-06-07 11:55:59 +08:00
sagitchu e569aedd3e feat(nftables): calculate traffic deltas 2026-06-07 11:55:59 +08:00
sagitchu 35080aea2d feat(flow): expose forward owner metadata 2026-06-07 11:55:59 +08:00
sagitchu 85e588ffe9 feat(nftables): persist counter state 2026-06-07 11:55:59 +08:00
sagitchu 03524f4a65 feat(nftables): collect counters over ssh 2026-06-07 11:55:59 +08:00
sagitchu 9e69e020ab feat(nftables): parse traffic counters 2026-06-07 11:55:59 +08:00
sagitchu 14bbd3907d feat(nftables): render traffic counters 2026-06-07 11:55:59 +08:00
sagitchu ca8d8e92ba docs: plan nftables traffic stats 2026-06-07 11:55:59 +08:00
sagitchu 079474fa06 docs: design nftables traffic stats 2026-06-07 11:55:59 +08:00
sagit 6e249a54f4 feat: add nftables forwarding mode (#516)
Adds nftables forwarding support for nodes, including frontend mode
selection, backend rule rendering, SSH-based rule reconciliation, and
online-state handling for nftables nodes.\n\nVerification:\n-
go-backend: go test ./...\n- vite-frontend: pnpm run build
2026-06-01 19:56:44 +08:00
sagitchu fb798a4532 chore(frontend): refresh pnpm lockfile 2026-06-01 19:54:28 +08:00
sagitchu a599f383f5 feat: add nftables forwarding mode 2026-06-01 19:47:26 +08:00
sagitchu 6bfa7f0166 docs: design nftables forwarding 2026-05-30 22:14:38 +08:00
sagit 2d0c993c90 fix: allow admin access to sensitive configs (#511) 2026-05-17 23:09:31 +08:00
64 changed files with 13190 additions and 557 deletions
+10 -1
View File
@@ -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,并使用 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
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">
&quot;{forwardToResetFlow?.name}&quot;
</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,407 @@
# nftables 纯转发设计
**日期**: 2026-05-30
**状态**: 待审核
**作者**: Codex
## 概述
为 FLVX 增加一种不依赖 agent 的纯转发能力:节点可选择 `nftables` 转发模式,面板通过 SSH 在节点机器上下发和维护 nftables 规则。
第一阶段只支持端口级 DNAT/SNAT 纯转发。它不是 GOST 隧道能力的替代品,也不支持链路、限速、流量统计、连接数限制、Proxy Protocol、best exit 或 agent 诊断。目标是提供一个可靠、可回滚、可重建的轻量转发路径。
## 背景
当前 FLVX 的转发模型由三部分组成:
- `node` 表描述节点,现有本地节点通过 agent WebSocket 接收运行时命令。
- `tunnel` 表描述入口、出口和链路类型,`type=1` 表示端口转发,`type=2` 表示隧道转发。
- `forward` 表描述用户规则、入口端口和目标地址,运行时通过 GOST service 下发到入口节点。
nftables 模式的核心差异是没有 agent,因此不能复用现有 WebSocket command 通道,也不能依赖 agent 上报在线状态、流量和诊断结果。面板必须成为唯一控制面,通过 SSH 把数据库中的期望状态同步到远端 nftables。
## 用户决策
- 创建或编辑节点时选择转发模式。
- 选择 nftables 转发后,不需要安装 agent。
- nftables 转发不支持隧道、流量控制等能力,只支持纯转发。
- 规则由面板端维护,并通过 SSH 下放到节点。
## 推荐方案
新增节点运行时模式:
| 模式 | 含义 |
|------|------|
| `agent` | 默认模式,保持现有 GOST agent 行为 |
| `nftables` | 面板通过 SSH 管理 nftables 规则 |
业务层继续复用现有 `tunnel` 和 `forward` 概念,但对 nftables 模式加严格能力边界:
- nftables 节点只能创建端口转发隧道。
- nftables 隧道不能配置出口节点或转发链。
- 同一个隧道的入口节点必须全部是同一种运行时模式。
- nftables 转发规则创建、更新、删除时,由后端同步 SSH 规则。
- 面板提供节点级“测试 SSH”“重建规则”“清理 FLVX 规则”操作。
## 非目标
- 不支持 `tunnel.type=2` 隧道转发。
- 不支持多跳链路、远程面板共享节点和 federation runtime。
- 不支持 GOST service 能力:限速、每 IP 限速、最大连接数、Proxy Protocol、策略负载均衡。
- 不支持 agent 流量统计、实时系统指标、节点升级、回退、agent 安装命令。
- 不在第一阶段支持 HA 漂移、自动探活切换或复杂负载均衡。
- 不改写用户机器上的非 FLVX nftables 规则。
## 数据模型
### node 表
新增字段:
| 字段 | 类型 | 默认 | 说明 |
|------|------|------|------|
| `forward_mode` | string | `agent` | `agent` 或 `nftables` |
Go 模型使用 SQLite/PostgreSQL 兼容 tag:
```go
ForwardMode string `gorm:"column:forward_mode;type:varchar(20);not null;default:'agent'"`
```
### node_ssh_config 表
新增表保存 nftables 节点 SSH 配置。SSH 凭据不放进 `node` 主表,避免普通节点列表过度暴露敏感字段。
| 字段 | 说明 |
|------|------|
| `id` | 主键 |
| `node_id` | 关联节点,唯一 |
| `host` | SSH 主机,默认可使用 node.server_ip |
| `port` | SSH 端口,默认 22 |
| `username` | SSH 用户 |
| `auth_type` | `password` 或 `private_key` |
| `password` | 加密后密码,可为空 |
| `private_key` | 加密后私钥,可为空 |
| `passphrase` | 加密后私钥口令,可为空 |
| `sudo_mode` | `none` / `sudo` |
| `created_time` | 创建时间 |
| `updated_time` | 更新时间 |
第一阶段可使用现有配置密钥派生或面板本地密钥做对称加密;如果项目尚无统一密钥管理,应至少避免在列表 API 返回完整凭据。
### nft_rule_binding 表
记录面板认为已经应用到节点的规则状态,用于更新、删除、重建和错误展示。
| 字段 | 说明 |
|------|------|
| `id` | 主键 |
| `forward_id` | 转发规则 ID |
| `node_id` | 下发节点 ID |
| `in_port` | 入口端口 |
| `protocols` | 第一阶段固定 `tcp,udp` |
| `target_addr` | 目标地址 |
| `bind_ip` | 可选监听 IP |
| `rule_hash` | 当前期望规则 hash |
| `status` | `pending` / `applied` / `error` |
| `last_error` | 最近错误 |
| `applied_time` | 最近成功应用时间 |
| `created_time` | 创建时间 |
| `updated_time` | 更新时间 |
绑定表不是最终事实来源。最终期望状态仍从 `forward`、`forward_port`、`tunnel` 和 `chain_tunnel` 推导,绑定表只记录应用结果。
## API 行为
### 节点创建和更新
`/node/create` 和 `/node/update` 新增入参:
```json
{
"forwardMode": "nftables",
"sshConfig": {
"host": "203.0.113.10",
"port": 22,
"username": "root",
"authType": "private_key",
"privateKey": "-----BEGIN OPENSSH PRIVATE KEY-----...",
"passphrase": "",
"sudoMode": "none"
}
}
```
规则:
- `forwardMode` 缺省时按 `agent`。
- `agent` 节点保留现有字段和行为。
- `nftables` 节点要求 SSH 配置完整。
- 从 `agent` 切到 `nftables` 前,若该节点已有 agent 隧道链路或转发规则,应拒绝并提示先迁移或删除。
- 从 `nftables` 切回 `agent` 前,若存在 nftables 规则,应拒绝并提示先清理或迁移。
### 隧道创建和更新
创建 nftables 隧道仍使用 `/tunnel/create`,但后端根据入口节点模式校验能力。
规则:
- 入口节点为 nftables 时,`type` 必须为 `1`。
- 不允许提交 `outNodeId` 或 `chainNodes`。
- 入口节点必须在线的现有校验不能直接套用到 nftables 节点;应改为 SSH 可用性校验或允许保存后手动测试。
- 同一隧道入口节点不能混用 `agent` 和 `nftables`。
- 更新隧道时不允许改变运行时模式;需要通过迁移规则到新隧道实现。
### 转发创建和更新
选择 nftables 隧道时,`/forward/create` 和 `/forward/update` 强制收窄字段:
- `speedId` 必须为空。
- `ipSpeedId` 必须为空。
- `maxConn` 和 `ipMaxConn` 必须为 0。
- `proxyProtocol` 必须为 0。
- 第一阶段 `remoteAddr` 只允许单目标 `host:port`。
- `strategy` 固定为 `fifo` 或忽略。
创建流程:
1. 校验权限、隧道状态、端口占用和 nftables 能力边界。
2. 在数据库创建 `forward` 和 `forward_port`。
3. 通过 nftables runtime 对关联入口节点执行同步。
4. 若同步失败,回滚数据库创建,返回 SSH/nftables 错误。
更新流程:
1. 保存旧 forward 和端口绑定。
2. 更新数据库。
3. 同步 nftables 规则。
4. 若同步失败,回滚数据库状态并尝试恢复旧规则。
删除流程:
1. 先删除远端 nftables 规则。
2. 成功后删除数据库。
3. 如果远端删除失败,普通删除返回错误;强制删除可删除数据库并保留 binding 错误记录,提示用户稍后清理。
## 后端组件
新增 package:
```text
go-backend/internal/runtime/nftables/
```
建议拆分:
| 组件 | 职责 |
|------|------|
| `Manager` | 对 handler 暴露 Apply/Delete/Reconcile/Test 方法 |
| `Planner` | 从数据库记录生成节点级期望规则 |
| `Renderer` | 把期望规则渲染为 nftables 脚本 |
| `SSHRunner` | 负责 SSH 连接、sudo 包装、命令执行和超时 |
| `Parser` | 解析目标地址、协议和错误信息 |
handler 不直接执行 SSH,也不拼 nft 脚本;handler 只做业务校验并调用 runtime manager。
## nftables 规则设计
FLVX 只维护自己的 table,避免触碰用户已有规则:
```nft
table inet flvx {
chain prerouting {
type nat hook prerouting priority dstnat; policy accept;
}
chain postrouting {
type nat hook postrouting priority srcnat; policy accept;
}
chain forward {
type filter hook forward priority filter; policy accept;
}
}
```
每条 forward 生成 TCP 和 UDP 规则:
```nft
tcp dport 12345 dnat to 198.51.100.20:443 comment "flvx forward:42 tcp"
udp dport 12345 dnat to 198.51.100.20:443 comment "flvx forward:42 udp"
```
第一阶段默认生成 masquerade:
```nft
masquerade comment "flvx masquerade"
```
原因是大多数纯 DNAT 场景需要回程可达;如果不做 SNAT,目标服务回包可能绕过转发节点导致连接失败。后续可增加高级开关允许用户关闭 masquerade。
### 原子同步策略
推荐节点级 reconcile,而不是逐条追加:
1. 从数据库查询该节点所有 nftables forward。
2. 生成完整 `table inet flvx` 脚本。
3. 通过 SSH 执行 `nft -f <tempfile>`。
4. 成功后更新所有相关 `nft_rule_binding` 状态和 hash。
这样可以避免局部更新导致规则漂移,也能让“重建规则”与创建/更新走同一条路径。
## SSH 执行策略
基础要求:
- 默认超时 10-15 秒。
- 支持密码和私钥认证。
- 支持 `sudo nft ...`。
- 执行前检查 `command -v nft`。
- 执行前检查 `nft --version`,错误时提示安装 nftables。
- 所有临时脚本写入 `/tmp/flvx-nft-<nonce>.nft`,执行后删除。
建议命令流程:
```sh
cat > /tmp/flvx-nft-xxxx.nft <<'EOF'
table inet flvx {
...
}
EOF
nft list table inet flvx >/dev/null 2>&1 && nft delete table inet flvx || true
nft -f /tmp/flvx-nft-xxxx.nft
rm -f /tmp/flvx-nft-xxxx.nft
```
如果目标 nft 版本支持 `destroy table`,也可以把删除动作放进脚本:
```nft
destroy table inet flvx
table inet flvx {
...
}
```
实现时应按目标 nft 版本兼容性选择 `destroy` 或 shell 中先检测 `nft list table inet flvx`。
## 前端体验
### 节点页
节点表单新增“转发模式”:
- `Agent 节点`:默认,现有表单不变。
- `nftables 节点`:显示 SSH 配置区块,隐藏 agent 安装相关提示。
nftables 节点列表操作:
- 测试 SSH
- 重建规则
- 清理 FLVX nftables 规则
隐藏或禁用:
- 安装命令
- 升级
- 回退
- agent 协议开关
- 实时 agent 指标入口
### 隧道页
隧道类型文案建议改为更明确的运行时说明:
- `Agent 端口转发`
- `Agent 隧道转发`
- `nftables 纯转发`
如果保持现有 `端口转发 / 隧道转发` 选择器,则在选择 nftables 入口节点后禁用隧道转发,并提示“不支持出口节点和转发链”。
### 转发页
选择 nftables 隧道后:
- 隐藏限速、每 IP 限速、最大连接数、Proxy Protocol。
- 目标地址输入提示“第一阶段仅支持单目标 host:port”。
- 创建/更新失败时显示远端 SSH 或 nftables 错误。
## 错误处理
- SSH 连接失败:返回“SSH 连接失败”,保留底层错误摘要。
- 认证失败:返回“SSH 认证失败,请检查用户名和凭据”。
- `nft` 不存在:返回“节点未安装 nftables”。
- nft 脚本失败:返回 nft stderr 摘要,并记录到 `nft_rule_binding.last_error`。
- 下发超时:标记 binding 为 `error`,允许用户重试“重建规则”。
- 数据库成功但远端失败时,创建/更新路径应回滚数据库;批量重建路径不回滚业务规则,只记录错误。
## 安全边界
- SSH 凭据只在创建/更新时接收,列表 API 不返回明文。
- 私钥和密码在数据库中加密保存。
- 后端日志不得打印完整私钥、密码或 passphrase。
- nft 脚本只由后端 renderer 生成,禁止直接拼接用户提交的自由文本。
- `remoteAddr` 必须严格解析为 host/IP + port,端口必须为 1-65535。
- `inPort` 仍复用现有端口占用校验。
- comment 中只放 forward ID 和协议,不放用户输入。
## 与现有功能的关系
- `node/install` 对 nftables 节点返回错误或前端隐藏入口。
- `node/check-status` 对 nftables 节点可返回 SSH 测试状态,而不是 agent 在线状态。
- `forward/batch-redeploy` 对 nftables 规则执行节点级 reconcile。
- `tunnel/batch-redeploy` 遇到 nftables 隧道时只重建相关 nftables 节点规则,不发送 GOST chain/service 命令。
- federation 导入/共享第一阶段不支持 nftables 节点。
- backup/import 应包含新增 node mode、SSH 配置和 binding 状态;导出时默认不导出 SSH 明文凭据。
## 测试计划
后端单元测试:
- nftables 节点不能创建隧道转发。
- nftables 隧道不能包含出口节点或转发链。
- agent 和 nftables 节点不能混在同一隧道。
- nftables forward 拒绝限速、连接限制和 Proxy Protocol。
- nftables forward 拒绝多目标 remoteAddr。
- renderer 为 TCP/UDP 生成稳定脚本和 comment。
- SSH runner 正确隐藏敏感信息并返回 stderr 摘要。
后端集成测试:
- 创建 nftables forward 时数据库和 binding 同步成功。
- runtime 下发失败时创建回滚。
- 更新失败时数据库和旧规则尽量恢复。
- 删除失败时普通删除返回错误,强制删除保留清理提示。
前端验证:
- 节点表单按转发模式切换字段。
- nftables 节点隐藏安装/升级/回退操作。
- 隧道表单阻止 nftables 隧道转发配置。
- 转发表单选择 nftables 隧道后隐藏不支持字段。
验证命令:
```bash
(cd go-backend && go test ./...)
(cd vite-frontend && pnpm run build)
```
## 实施顺序
1. 数据模型和 repository:新增字段、SSH 配置表、binding 表和查询方法。
2. nftables runtime:实现 planner、renderer、SSH runner、manager。
3. handler 校验:节点、隧道、转发 create/update/delete 接入 runtime。
4. 前端节点表单:增加转发模式和 SSH 配置。
5. 前端隧道/转发表单:按 nftables 能力收窄 UI。
6. 批量重建和清理操作:提供运维入口。
7. 测试与文案打磨。
## 第一阶段固定决策
本设计先固定以下选择,除非审核时调整:
- 第一阶段同时下发 TCP 和 UDP。
- 第一阶段只支持单目标。
- 第一阶段默认启用 masquerade。
- nftables 节点的“在线状态”以 SSH 测试为准,而不是常驻连接。
@@ -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`
不需要数据库迁移或新增依赖。
@@ -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
}
@@ -240,6 +242,19 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
if err != nil {
return nil, err
}
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
if err != nil {
return nil, err
}
if nftMode {
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
return nil, err
}
if len(entryNodeIDs) == 0 {
return nil, errors.New("nftables 转发缺少入口节点")
}
return nil, h.syncNftablesNode(entryNodeIDs[0])
}
ports, err := h.listForwardPorts(forward.ID)
if err != nil {
return nil, err
@@ -759,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
@@ -774,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
@@ -1474,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 {
@@ -1504,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",
})
@@ -1744,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"
}
@@ -1788,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{}{
@@ -1806,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)
}
@@ -1828,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 == "" {
@@ -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)
}
}
+26 -5
View File
@@ -23,6 +23,7 @@ import (
"go-backend/internal/license"
"go-backend/internal/metrics"
"go-backend/internal/monitoring"
runtimenft "go-backend/internal/runtime/nftables"
"go-backend/internal/security"
"go-backend/internal/store/repo"
"go-backend/internal/ws"
@@ -31,11 +32,12 @@ import (
)
type Handler struct {
repo *repo.Repository
jwtSecret string
wsServer *ws.Server
metrics *metrics.IngestionService
healthCheck *health.Checker
repo *repo.Repository
jwtSecret string
wsServer *ws.Server
metrics *metrics.IngestionService
healthCheck *health.Checker
nftablesManager nftablesRuntimeManager
captchaMu sync.Mutex
captchaTokens map[string]int64
@@ -108,6 +110,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
wsServer: ws.NewServer(repo, jwtSecret),
metrics: metrics.NewIngestionService(repo),
healthCheck: nil,
nftablesManager: runtimenft.NewManager(nil),
captchaTokens: make(map[string]int64),
pendingUpgradeRedeploy: make(map[int64]struct{}),
nodeOnlineRedeployAt: make(map[int64]time.Time),
@@ -190,6 +193,9 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/node/batch-upgrade", h.nodeBatchUpgrade)
mux.HandleFunc("/api/v1/node/rollback", h.nodeRollback)
mux.HandleFunc("/api/v1/node/releases", h.listReleases)
mux.HandleFunc("/api/v1/node/nftables/test", h.nodeNftablesTest)
mux.HandleFunc("/api/v1/node/nftables/reconcile", h.nodeNftablesReconcile)
mux.HandleFunc("/api/v1/node/nftables/clear", h.nodeNftablesClear)
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
@@ -215,6 +221,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)
@@ -1008,6 +1015,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())
@@ -1055,6 +1063,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())
}
@@ -1103,11 +1112,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
+54 -2
View File
@@ -2,11 +2,14 @@ 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 +23,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,6 +33,7 @@ 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) {
@@ -64,7 +68,7 @@ func (h *Handler) validateLicenseJob() {
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
client := license.NewKeygenClient(accountID, "")
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
if err != nil {
// Network error or timeout. Grace period by not revoking immediately here.
return
@@ -128,6 +132,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()
+326 -10
View File
@@ -341,6 +341,11 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
now := time.Now().UnixMilli()
inx := h.repo.NextIndex("node")
forwardMode := defaultNodeForwardMode(asString(req["forwardMode"]))
status := 0
if forwardMode == "nftables" {
status = 1
}
if err := h.repo.CreateNode(
name,
randomToken(16),
@@ -357,7 +362,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
asInt(req["tls"], 0),
asInt(req["socks"], 0),
now,
0,
status,
defaultString(asString(req["tcpListenAddr"]), "[::]"),
defaultString(asString(req["udpListenAddr"]), "[::]"),
inx,
@@ -366,10 +371,20 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
nullableText(asString(req["remoteToken"])),
nullableText(asString(req["remoteConfig"])),
nullableText(asString(req["extraIPs"])),
forwardMode,
); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
nodeID, err := h.findCreatedNodeID(name, serverIP)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.persistNodeSSHConfig(nodeID, req, forwardMode, now, false); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -402,6 +417,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 {
@@ -409,7 +430,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
@@ -428,6 +450,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
nullableText(strings.TrimSpace(asString(req["remark"]))),
nullableUnixMilli(asInt64(req["expiryTime"], 0)),
nullableText(normalizeNodeRenewalCycle(asString(req["renewalCycle"]))),
forwardMode,
newHTTP,
newTLS,
newSocks,
@@ -438,9 +461,142 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.persistNodeSSHConfig(id, req, forwardMode, now, true); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if forwardMode == "nftables" && currentStatus != 1 {
if err := h.repo.UpdateNodeStatus(id, 1); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) findCreatedNodeID(name, serverIP string) (int64, error) {
if h == nil || h.repo == nil {
return 0, errors.New("handler not initialized")
}
nodes, err := h.repo.ListNodes()
if err != nil {
return 0, err
}
for i := len(nodes) - 1; i >= 0; i-- {
item := nodes[i]
if asString(item["name"]) != name {
continue
}
if asString(item["serverIp"]) != serverIP {
continue
}
if nodeID := asInt64(item["id"], 0); nodeID > 0 {
return nodeID, nil
}
}
return 0, errors.New("节点创建成功,但未能查询到节点记录")
}
func (h *Handler) persistNodeSSHConfig(nodeID int64, req map[string]interface{}, forwardMode string, now int64, preserveSecrets bool) error {
if h == nil || h.repo == nil {
return errors.New("handler not initialized")
}
if nodeID <= 0 {
return errors.New("节点ID不能为空")
}
if forwardMode != "nftables" {
return h.repo.DeleteNodeSSHConfig(nodeID)
}
cfgMap := asMap(req["sshConfig"])
host := strings.TrimSpace(asString(cfgMap["host"]))
if host == "" {
host = strings.TrimSpace(asString(req["serverIp"]))
}
port := asInt(cfgMap["port"], 22)
username := strings.TrimSpace(asString(cfgMap["username"]))
authType := strings.TrimSpace(asString(cfgMap["authType"]))
password := asString(cfgMap["password"])
privateKey := asString(cfgMap["privateKey"])
passphrase := asString(cfgMap["passphrase"])
sudoMode := strings.TrimSpace(asString(cfgMap["sudoMode"]))
if preserveSecrets {
existing, err := h.repo.GetNodeSSHConfig(nodeID)
if err != nil && !errors.Is(err, sql.ErrNoRows) {
return err
}
if existing != nil {
if host == "" {
host = strings.TrimSpace(existing.Host)
}
if port <= 0 {
port = existing.Port
}
if username == "" {
username = strings.TrimSpace(existing.Username)
}
if authType == "" {
authType = strings.TrimSpace(existing.AuthType)
}
if strings.TrimSpace(password) == "" && existing.Password.Valid {
password = existing.Password.String
}
if strings.TrimSpace(privateKey) == "" && existing.PrivateKey.Valid {
privateKey = existing.PrivateKey.String
}
if strings.TrimSpace(passphrase) == "" && existing.Passphrase.Valid {
passphrase = existing.Passphrase.String
}
if sudoMode == "" {
sudoMode = strings.TrimSpace(existing.SudoMode)
}
}
}
if host == "" || username == "" {
return errors.New("nftables 节点 SSH 配置不完整")
}
if port <= 0 || port > 65535 {
return errors.New("nftables 节点 SSH 端口无效")
}
authType = strings.ToLower(authType)
switch authType {
case "password":
if strings.TrimSpace(password) == "" {
return errors.New("nftables 节点 SSH 密码不能为空")
}
privateKey = ""
case "private_key", "":
authType = "private_key"
if strings.TrimSpace(privateKey) == "" {
return errors.New("nftables 节点 SSH 私钥不能为空")
}
password = ""
default:
return errors.New("nftables 节点 SSH 认证方式无效")
}
switch strings.ToLower(sudoMode) {
case "", "none":
sudoMode = "none"
case "sudo", "sudo_su":
default:
return errors.New("nftables 节点 sudo 模式无效")
}
return h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
Host: host,
Port: port,
Username: username,
AuthType: authType,
Password: password,
PrivateKey: privateKey,
Passphrase: passphrase,
SudoMode: sudoMode,
}, now)
}
func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
@@ -648,6 +804,16 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
if strings.TrimSpace(inIP) == "" {
inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
}
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
for _, inNode := range runtimeState.InNodes {
if inNode.NodeID > 0 {
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
}
}
if err := h.validateNftablesTunnelStateTx(tx, entryNodeIDs); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
if len(runtimeState.InNodes) > 0 {
firstNodeID := runtimeState.InNodes[0].NodeID
@@ -934,6 +1100,16 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
}
runtimeState.TunnelID = id
runtimeState.IPPreference = ipPreference
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
for _, inNode := range runtimeState.InNodes {
if inNode.NodeID > 0 {
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
}
}
if err := h.validateNftablesTunnelState(entryNodeIDs); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
@@ -1700,6 +1876,24 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
failures = appendBatchFailure(failures, tunnelID, tunnelName, tunnelErr)
continue
}
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(tunnelID); modeErr != nil {
fail++
failures = appendBatchFailure(failures, tunnelID, tunnelName, modeErr)
continue
} else if nftMode {
if len(entryNodeIDs) == 0 {
fail++
failures = appendBatchFailureReason(failures, tunnelID, tunnelName, "nftables 转发缺少入口节点")
continue
}
if reconcileErr := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); reconcileErr != nil {
fail++
failures = appendBatchFailure(failures, tunnelID, tunnelName, reconcileErr)
continue
}
success++
continue
}
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
fail++
failures = appendBatchFailure(failures, tunnelID, tunnelName, err)
@@ -1909,6 +2103,17 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
port = 10000
}
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
isNftTunnel, _, err := h.tunnelUsesNftables(tunnelID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if isNftTunnel {
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
inIp := strings.TrimSpace(asString(req["inIp"]))
if inIp != "" && len(entryNodes) > 1 {
response.WriteJSON(w, response.ErrDefault("多入口隧道的转发不支持自定义监听IP"))
@@ -1944,8 +2149,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
@@ -2016,6 +2223,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
if remoteAddr == "" {
remoteAddr = forward.RemoteAddr
}
isNftTunnel, entryNodes, err := h.tunnelUsesNftables(tunnelID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if isNftTunnel {
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
if actorRole != 0 && !h.allowLocalRemoteAddr() {
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
response.WriteJSON(w, response.Err(403, err.Error()))
@@ -2123,8 +2341,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
}
@@ -2202,7 +2422,17 @@ 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 {
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 {
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
}
@@ -2210,6 +2440,12 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
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())
}
@@ -2218,7 +2454,7 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
if id <= 0 {
return
}
_, _, _, err := h.resolveForwardAccess(r, id)
forward, _, _, err := h.resolveForwardAccess(r, id)
if err != nil {
if errors.Is(err, errForwardNotFound) {
response.WriteJSON(w, response.ErrDefault("转发不存在"))
@@ -2234,6 +2470,13 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.repo.DeleteNftRuleBindingsByForward(id)
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr == nil && nftMode && len(entryNodeIDs) > 0 {
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -2259,6 +2502,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 {
@@ -2351,7 +2614,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
@@ -2359,9 +2634,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}))
}
@@ -2462,6 +2744,24 @@ func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
failures = appendBatchFailure(failures, id, "", accessErr)
continue
}
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
}
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
f++
failures = appendBatchFailure(failures, id, forward.Name, err)
} else {
s++
}
continue
}
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
f++
failures = appendBatchFailure(failures, id, forward.Name, err)
@@ -4452,7 +4752,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(),
)
@@ -4705,6 +5005,13 @@ func asMapSlice(v interface{}) []map[string]interface{} {
return out
}
func asMap(v interface{}) map[string]interface{} {
if m, ok := v.(map[string]interface{}); ok && m != nil {
return m
}
return map[string]interface{}{}
}
func asString(v interface{}) string {
switch t := v.(type) {
case nil:
@@ -4843,6 +5150,15 @@ func normalizeNodeRenewalCycle(v string) string {
}
}
func defaultNodeForwardMode(mode string) string {
switch strings.TrimSpace(strings.ToLower(mode)) {
case "nftables":
return "nftables"
default:
return "agent"
}
}
func nullableInt(v *int64) interface{} {
if v == nil {
return nil
@@ -0,0 +1,371 @@
package handler
import (
"context"
"database/sql"
"errors"
"fmt"
"net"
"net/http"
"strings"
"time"
"go-backend/internal/http/response"
runtimenft "go-backend/internal/runtime/nftables"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
"gorm.io/gorm"
)
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 {
return strings.EqualFold(strings.TrimSpace(mode), runtimenft.ModeNftables)
}
func (h *Handler) nodeUsesNftables(nodeID int64) (bool, error) {
return h.nodeUsesNftablesTx(nil, nodeID)
}
func (h *Handler) nodeUsesNftablesTx(tx *gorm.DB, nodeID int64) (bool, error) {
if h == nil || h.repo == nil {
return false, errors.New("handler not initialized")
}
var (
mode string
err error
)
if tx != nil {
mode, err = h.repo.GetNodeForwardModeTx(tx, nodeID)
} else {
mode, err = h.repo.GetNodeForwardMode(nodeID)
}
if err != nil {
return false, err
}
return isNftablesForwardMode(mode), nil
}
func (h *Handler) tunnelUsesNftables(tunnelID int64) (bool, []int64, error) {
entryNodeIDs, err := h.tunnelEntryNodeIDs(tunnelID)
if err != nil {
return false, nil, err
}
for _, nodeID := range entryNodeIDs {
ok, modeErr := h.nodeUsesNftables(nodeID)
if modeErr != nil {
return false, nil, modeErr
}
if ok {
return true, entryNodeIDs, nil
}
}
return false, entryNodeIDs, nil
}
func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAddr string, entryNodeIDs []int64) error {
if tunnel == nil {
return errors.New("隧道不存在")
}
if tunnel.Type != 1 {
return errors.New("nftables 节点仅支持直连隧道")
}
if len(entryNodeIDs) != 1 {
return errors.New("nftables 节点仅支持单入口隧道")
}
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
}
func sshConfigFromModel(cfg *model.NodeSSHConfig) (runtimenft.SSHConfig, error) {
if cfg == nil {
return runtimenft.SSHConfig{}, errors.New("节点缺少 SSH 配置")
}
if strings.TrimSpace(cfg.Host) == "" || strings.TrimSpace(cfg.Username) == "" {
return runtimenft.SSHConfig{}, errors.New("节点 SSH 配置不完整")
}
return runtimenft.SSHConfig{
Host: strings.TrimSpace(cfg.Host),
Port: cfg.Port,
Username: strings.TrimSpace(cfg.Username),
AuthType: strings.TrimSpace(cfg.AuthType),
Password: cfg.Password.String,
PrivateKey: cfg.PrivateKey.String,
Passphrase: cfg.Passphrase.String,
SudoMode: strings.TrimSpace(cfg.SudoMode),
}, nil
}
func (h *Handler) validateNftablesTunnelState(entryNodeIDs []int64) error {
return h.validateNftablesTunnelStateTx(nil, entryNodeIDs)
}
func (h *Handler) validateNftablesTunnelStateTx(tx *gorm.DB, entryNodeIDs []int64) error {
if h == nil || h.repo == nil {
return errors.New("handler not initialized")
}
for _, nodeID := range entryNodeIDs {
isNft, err := h.nodeUsesNftablesTx(tx, nodeID)
if err != nil {
return err
}
if !isNft {
continue
}
var cfg *model.NodeSSHConfig
if tx != nil {
cfg, err = h.repo.GetNodeSSHConfigTx(tx, nodeID)
} else {
cfg, err = h.repo.GetNodeSSHConfig(nodeID)
}
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return errors.New("nftables 节点缺少 SSH 配置")
}
return err
}
sshCfg, err := sshConfigFromModel(cfg)
if err != nil {
return err
}
if h.nftablesManager == nil {
return errors.New("nftables manager not initialized")
}
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
return fmt.Errorf("nftables 节点能力校验失败: %w", err)
}
}
return nil
}
func (h *Handler) buildNftablesNodePlan(nodeID int64) (runtimenft.NodePlan, *model.NodeSSHConfig, error) {
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
if err != nil {
return runtimenft.NodePlan{}, nil, err
}
forwards, err := h.repo.ListActiveForwardsByNode(nodeID)
if err != nil {
return runtimenft.NodePlan{}, nil, err
}
plan := runtimenft.NodePlan{NodeID: nodeID, Rules: make([]runtimenft.Rule, 0, len(forwards))}
for i := range forwards {
forward := &forwards[i]
tunnel, err := h.getTunnelRecord(forward.TunnelID)
if err != nil || tunnel == nil || tunnel.Status != 1 {
continue
}
entryNodeIDs, err := h.tunnelEntryNodeIDs(forward.TunnelID)
if err != nil {
return runtimenft.NodePlan{}, nil, err
}
if len(entryNodeIDs) != 1 || entryNodeIDs[0] != nodeID {
continue
}
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
return runtimenft.NodePlan{}, nil, err
}
ports, err := h.listForwardPorts(forward.ID)
if err != nil {
return runtimenft.NodePlan{}, nil, err
}
for _, fp := range ports {
if fp.NodeID != nodeID {
continue
}
target, err := runtimenft.ParseSingleTarget(forward.RemoteAddr)
if err != nil {
return runtimenft.NodePlan{}, nil, err
}
plan.Rules = append(plan.Rules, runtimenft.Rule{
ForwardID: forward.ID,
InPort: fp.Port,
BindIP: strings.TrimSpace(fp.InIP),
TargetHost: target.Host,
TargetPort: target.Port,
Protocols: []string{"tcp", "udp"},
})
}
}
return plan, cfg, nil
}
func (h *Handler) syncNftablesNode(nodeID int64) error {
if h == nil || h.repo == nil {
return errors.New("handler not initialized")
}
plan, cfgModel, err := h.buildNftablesNodePlan(nodeID)
if err != nil {
return err
}
sshCfg, err := sshConfigFromModel(cfgModel)
if err != nil {
return err
}
result, err := h.nftablesManager.Reconcile(context.Background(), sshCfg, plan)
now := time.Now().UnixMilli()
if err != nil {
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
for _, binding := range bindings {
_ = h.repo.MarkNftRuleBindingError(binding.ForwardID, nodeID, err.Error(), now)
}
return err
}
activeForwardIDs := make(map[int64]struct{}, len(plan.Rules))
for _, rule := range plan.Rules {
activeForwardIDs[rule.ForwardID] = struct{}{}
hash := result.Hashes[rule.ForwardID]
_ = h.repo.UpsertNftRuleBinding(modelToRuleBindingInput(nodeID, rule, hash), now)
}
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
for _, binding := range bindings {
if _, ok := activeForwardIDs[binding.ForwardID]; ok {
continue
}
_ = h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID)
}
return nil
}
func modelToRuleBindingInput(nodeID int64, rule runtimenft.Rule, hash string) repo.NftRuleBindingInput {
return repo.NftRuleBindingInput{
ForwardID: rule.ForwardID,
NodeID: nodeID,
InPort: rule.InPort,
Protocols: strings.Join(rule.Protocols, ","),
TargetAddr: fmt.Sprintf("%s:%d", rule.TargetHost, rule.TargetPort),
BindIP: rule.BindIP,
RuleHash: hash,
Status: runtimenft.StatusApplied,
}
}
func (h *Handler) nftablesNodeIDFromRequest(r *http.Request, w http.ResponseWriter) (int64, bool) {
nodeID := asInt64FromBodyKey(r, w, "nodeId")
if nodeID <= 0 {
return 0, false
}
return nodeID, true
}
func (h *Handler) loadNftablesSSHConfig(nodeID int64) (runtimenft.SSHConfig, error) {
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return runtimenft.SSHConfig{}, errors.New("nftables 节点缺少 SSH 配置")
}
return runtimenft.SSHConfig{}, err
}
return sshConfigFromModel(cfg)
}
func (h *Handler) clearNftablesNode(nodeID int64) error {
if h == nil || h.repo == nil {
return errors.New("handler not initialized")
}
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
if err != nil {
return err
}
if h.nftablesManager == nil {
return errors.New("nftables manager not initialized")
}
if err := h.nftablesManager.Clear(context.Background(), sshCfg); err != nil {
return err
}
bindings, listErr := h.repo.ListNftRuleBindingsByNode(nodeID)
if listErr != nil {
return listErr
}
for _, binding := range bindings {
if err := h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID); err != nil {
return err
}
}
return nil
}
func (h *Handler) reconcileNftablesNodeByRequest(nodeID int64) error {
usesNft, err := h.nodeUsesNftables(nodeID)
if err != nil {
return err
}
if !usesNft {
return errors.New("节点未启用 nftables 转发模式")
}
return h.syncNftablesNode(nodeID)
}
func (h *Handler) nodeNftablesTest(w http.ResponseWriter, r *http.Request) {
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
if !ok {
return
}
usesNft, err := h.nodeUsesNftables(nodeID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if !usesNft {
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
return
}
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
if h.nftablesManager == nil {
response.WriteJSON(w, response.Err(-2, "nftables manager not initialized"))
return
}
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeNftablesReconcile(w http.ResponseWriter, r *http.Request) {
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
if !ok {
return
}
if err := h.reconcileNftablesNodeByRequest(nodeID); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) nodeNftablesClear(w http.ResponseWriter, r *http.Request) {
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
if !ok {
return
}
usesNft, err := h.nodeUsesNftables(nodeID)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if !usesNft {
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
return
}
if err := h.clearNftablesNode(nodeID); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -0,0 +1,711 @@
package handler
import (
"bytes"
"context"
"database/sql"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"sync"
"testing"
"time"
"go-backend/internal/auth"
"go-backend/internal/http/middleware"
runtimenft "go-backend/internal/runtime/nftables"
"go-backend/internal/store/repo"
)
type fakeNftablesManager struct {
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
if f.reconcileErr != nil {
return runtimenft.ApplyResult{}, f.reconcileErr
}
return runtimenft.ApplyResult{
NodeID: plan.NodeID,
Script: "table inet flvx {}",
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
}
func TestTunnelCreateRejectsNftablesEntryNodeWithoutSSHConfig(t *testing.T) {
fixture := setupNftablesHandler(t)
err := fixture.handler.validateNftablesTunnelState([]int64{fixture.nodeID})
if err == nil {
t.Fatalf("expected validation failure")
}
if !strings.Contains(err.Error(), "SSH") {
t.Fatalf("expected SSH config validation error, got %q", err)
}
}
func TestTunnelUpdateRejectsNftablesEntryNodeWhenCapabilityTestFails(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
manager := &fakeNftablesManager{testErr: errors.New("ssh failed")}
h.nftablesManager = manager
seedNftablesSSHConfig(t, h, fixture.nodeID)
err := h.validateNftablesTunnelState([]int64{fixture.nodeID})
if err == nil {
t.Fatalf("expected validation failure")
}
if !strings.Contains(err.Error(), "ssh failed") {
t.Fatalf("expected capability error in response, got %q", err)
}
}
func TestSyncForwardServicesWithWarningsUsesNftablesRuntime(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
manager := &fakeNftablesManager{}
h.nftablesManager = manager
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")
warnings, err := h.syncForwardServicesWithWarnings(forward, "UpdateService", true)
if err != nil {
t.Fatalf("sync forward services: %v", err)
}
if len(warnings) != 0 {
t.Fatalf("expected no warnings, got %v", warnings)
}
if manager.reconcileHit != 1 {
t.Fatalf("expected nftables reconcile to run once, got %d", manager.reconcileHit)
}
if manager.lastPlan.NodeID != fixture.nodeID {
t.Fatalf("expected plan for node %d, got %+v", fixture.nodeID, manager.lastPlan)
}
if len(manager.lastPlan.Rules) != 1 || manager.lastPlan.Rules[0].ForwardID != forward.ID {
t.Fatalf("unexpected plan: %+v", manager.lastPlan)
}
}
func TestNodeNftablesTestEndpointRunsCapabilityCheck(t *testing.T) {
fixture := setupNftablesHandler(t)
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
manager := &fakeNftablesManager{}
fixture.handler.nftablesManager = manager
res := postJSONToHandler(t, fixture.handler.nodeNftablesTest, map[string]int64{"nodeId": fixture.nodeID})
assertNftablesSuccess(t, res)
if manager.lastConfig.Host != "203.0.113.10" {
t.Fatalf("expected SSH config to be passed to manager, got %+v", manager.lastConfig)
}
}
func TestNodeNftablesReconcileEndpointPersistsBindings(t *testing.T) {
fixture := setupNftablesHandler(t)
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
tunnelID := seedTunnelForNftables(t, fixture.handler, "nft-tunnel", fixture.nodeID)
forward := seedForwardForNftables(t, fixture.handler, tunnelID, fixture.nodeID, "203.0.113.9:8080")
manager := &fakeNftablesManager{}
fixture.handler.nftablesManager = manager
res := postJSONToHandler(t, fixture.handler.nodeNftablesReconcile, map[string]int64{"nodeId": fixture.nodeID})
assertNftablesSuccess(t, res)
if manager.reconcileHit != 1 {
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
}
bindings, err := fixture.handler.repo.ListNftRuleBindingsByNode(fixture.nodeID)
if err != nil {
t.Fatalf("list bindings: %v", err)
}
if len(bindings) != 1 || bindings[0].ForwardID != forward.ID {
t.Fatalf("unexpected bindings: %+v", bindings)
}
}
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)
now := time.Now().UnixMilli()
if err := fixture.handler.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
ForwardID: 99,
NodeID: fixture.nodeID,
InPort: 24000,
Protocols: "tcp",
TargetAddr: "203.0.113.9:8080",
Status: runtimenft.StatusApplied,
}, now); err != nil {
t.Fatalf("seed binding: %v", err)
}
manager := &fakeNftablesManager{}
fixture.handler.nftablesManager = manager
res := postJSONToHandler(t, fixture.handler.nodeNftablesClear, map[string]int64{"nodeId": fixture.nodeID})
assertNftablesSuccess(t, res)
if manager.clearHit != 1 {
t.Fatalf("expected clear once, got %d", manager.clearHit)
}
if bindings, err := fixture.handler.repo.ListNftRuleBindingsByNode(fixture.nodeID); err != nil {
t.Fatalf("list bindings after clear: %v", err)
} else if len(bindings) != 0 {
t.Fatalf("expected bindings to be cleared, got %+v", bindings)
}
}
func TestNodeCreatePersistsNftablesSSHConfig(t *testing.T) {
fixture := setupNftablesHandler(t)
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
"name": "nft-node-created",
"serverIp": "203.0.113.20",
"serverIpV4": "203.0.113.20",
"port": "20000-20100",
"forwardMode": "nftables",
"sshConfig": map[string]interface{}{
"host": "203.0.113.21",
"port": 2222,
"username": "root",
"authType": "private_key",
"privateKey": "TEST-PRIVATE-KEY",
"passphrase": "secret",
"sudoMode": "sudo",
},
})
res := httptest.NewRecorder()
fixture.handler.nodeCreate(res, req)
assertNftablesSuccessWithBody(t, res)
nodes, err := fixture.handler.repo.ListNodes()
if err != nil {
t.Fatalf("list nodes: %v", err)
}
var createdNodeID int64
for _, item := range nodes {
if item["name"] == "nft-node-created" {
createdNodeID = item["id"].(int64)
break
}
}
if createdNodeID <= 0 {
t.Fatalf("expected created node to exist")
}
createdNode, err := fixture.handler.repo.GetNodeRecord(createdNodeID)
if err != nil {
t.Fatalf("load created node: %v", err)
}
if createdNode == nil {
t.Fatal("expected created node record, got nil")
}
if createdNode.Status != 1 {
t.Fatalf("expected nftables node to be online, got status %d", createdNode.Status)
}
cfg, err := fixture.handler.repo.GetNodeSSHConfig(createdNodeID)
if err != nil {
t.Fatalf("load ssh config: %v", err)
}
if cfg.Host != "203.0.113.21" || cfg.Port != 2222 || cfg.Username != "root" || cfg.AuthType != "private_key" {
t.Fatalf("unexpected ssh config: %+v", cfg)
}
if !cfg.PrivateKey.Valid || cfg.PrivateKey.String != "TEST-PRIVATE-KEY" {
t.Fatalf("expected private key to persist, got %+v", cfg)
}
}
func TestNodeUpdatePreservesExistingNftablesSecretsWhenFieldsOmitted(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",
"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.Host != "203.0.113.30" || cfg.Username != "admin" || cfg.AuthType != "password" {
t.Fatalf("unexpected ssh config after update: %+v", cfg)
}
if !cfg.Password.Valid || cfg.Password.String != "secret" {
t.Fatalf("expected password secret to be preserved, got %+v", cfg)
}
}
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
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")
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
ForwardID: forward.ID,
NodeID: fixture.nodeID,
InPort: 20000,
Protocols: "tcp,udp",
TargetAddr: "203.0.113.9:8080",
Status: runtimenft.StatusApplied,
}, time.Now().UnixMilli()); err != nil {
t.Fatalf("seed binding: %v", err)
}
manager := &fakeNftablesManager{}
h.nftablesManager = manager
req := newAuthenticatedJSONRequest(t, map[string]int64{"id": forward.ID})
req.URL.Path = "/api/v1/forward/force-delete"
res := httptest.NewRecorder()
mux := http.NewServeMux()
h.Register(mux)
mux.ServeHTTP(res, req)
assertNftablesSuccessWithBody(t, res)
if manager.reconcileHit != 1 {
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
}
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
t.Fatalf("expected forward to be deleted, got %v", err)
}
if bindings, err := h.repo.ListNftRuleBindingsByNode(fixture.nodeID); err != nil {
t.Fatalf("list bindings after delete: %v", err)
} else if len(bindings) != 0 {
t.Fatalf("expected no bindings after delete, got %+v", bindings)
}
}
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
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.forwardBatchRedeploy(res, req)
assertNftablesSuccessWithBody(t, res)
if manager.reconcileHit != 1 {
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
}
}
func TestTunnelBatchRedeployUsesNftablesReconcile(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
seedNftablesSSHConfig(t, h, fixture.nodeID)
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
manager := &fakeNftablesManager{}
h.nftablesManager = manager
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {tunnelID}})
res := httptest.NewRecorder()
h.tunnelBatchRedeploy(res, req)
assertNftablesSuccessWithBody(t, res)
if manager.reconcileHit != 1 {
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
}
}
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()
dbPath := filepath.Join(t.TempDir(), "handler-nftables.sqlite")
r, err := repo.Open(dbPath)
if err != nil {
t.Fatalf("open repo: %v", err)
}
h := New(r, "test-secret")
now := time.Now().UnixMilli()
if _, err := r.CreateUser("admin", "hash", 0, now+86400000, 1, 1, 100, 1, 0, now); err != nil {
t.Fatalf("create user: %v", err)
}
if err := r.CreateNode("nft-node", "secret", "198.51.100.10", nil, nil, "1000-65535", nil, nil, nil, nil, nil, 0, 0, 0, now, 1, "", "", 1, 0, nil, nil, nil, nil, "nftables"); err != nil {
t.Fatalf("create node: %v", err)
}
node, err := r.GetNodeRecord(1)
if err != nil || node == nil {
t.Fatalf("get node: %v", err)
}
return nftablesTestFixture{handler: h, nodeID: node.ID}
}
func seedNftablesSSHConfig(t *testing.T, h *Handler, nodeID int64) {
t.Helper()
if err := h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
Host: "203.0.113.10",
Port: 22,
Username: "root",
AuthType: "password",
Password: "secret",
SudoMode: "none",
}, time.Now().UnixMilli()); err != nil {
t.Fatalf("upsert ssh config: %v", err)
}
}
func seedTunnelForNftables(t *testing.T, h *Handler, name string, nodeID int64) int64 {
t.Helper()
now := time.Now().UnixMilli()
tx := h.repo.BeginTx()
if tx == nil {
t.Fatal("begin tx: nil transaction")
}
if tx.Error != nil {
t.Fatalf("begin tx: %v", tx.Error)
}
tunnelID, err := h.repo.CreateTunnelTx(tx, name, 1, 1, 1, now, 1, nil, 1, "", "", 0)
if err != nil {
_ = tx.Rollback().Error
t.Fatalf("create tunnel: %v", err)
}
if err := h.repo.CreateChainTunnelTx(tx, tunnelID, "1", nodeID, sql.NullInt64{}, "", 1, "tls", ""); err != nil {
_ = tx.Rollback().Error
t.Fatalf("create chain tunnel: %v", err)
}
if err := tx.Commit().Error; err != nil {
_ = tx.Rollback().Error
t.Fatalf("commit tx: %v", err)
}
return tunnelID
}
func seedForwardForNftables(t *testing.T, h *Handler, tunnelID, nodeID int64, remoteAddr string) *forwardRecord {
t.Helper()
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, 0, 0,
)
if err != nil {
t.Fatalf("create forward: %v", err)
}
forward, err := h.getForwardRecord(forwardID)
if err != nil {
t.Fatalf("get forward: %v", err)
}
return forward
}
func postJSONToHandler(t *testing.T, fn func(http.ResponseWriter, *http.Request), payload any) *httptest.ResponseRecorder {
t.Helper()
body, err := json.Marshal(payload)
if err != nil {
t.Fatalf("marshal payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
res := httptest.NewRecorder()
fn(res, req)
return res
}
func newAuthenticatedJSONRequest(t *testing.T, payload any) *http.Request {
t.Helper()
body, err := json.Marshal(payload)
if err != nil {
t.Fatalf("marshal payload: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
token, err := auth.GenerateToken(1, "admin", 0, "test-secret")
if err != nil {
t.Fatalf("create token: %v", err)
}
req.Header.Set("Authorization", token)
claims, ok := auth.ValidateToken(token, "test-secret")
if !ok {
t.Fatalf("validate token failed")
}
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
}
func assertNftablesSuccess(t *testing.T, res *httptest.ResponseRecorder) {
t.Helper()
assertNftablesSuccessWithBody(t, res)
}
func assertNftablesSuccessWithBody(t *testing.T, res *httptest.ResponseRecorder) {
t.Helper()
var payload struct {
Code int `json:"code"`
Msg string `json:"msg"`
}
if res.Code != http.StatusOK {
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
}
if err := json.NewDecoder(res.Body).Decode(&payload); err != nil {
t.Fatalf("decode response: %v", err)
}
if payload.Code != 0 {
t.Fatalf("expected success, got %+v", payload)
}
}
@@ -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
}
@@ -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)
}
}
}
@@ -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")
}
}
@@ -0,0 +1,65 @@
package nftables
import (
"context"
"errors"
)
type Manager struct {
runner Runner
}
func NewManager(runner Runner) *Manager {
if runner == nil {
runner = NewSSHRunner()
}
return &Manager{runner: runner}
}
func (m *Manager) Test(ctx context.Context, cfg SSHConfig) error {
if err := m.ensureInitialized(); err != nil {
return err
}
return m.runner.Test(ctx, cfg)
}
func (m *Manager) Reconcile(ctx context.Context, cfg SSHConfig, plan NodePlan) (ApplyResult, error) {
if err := m.ensureInitialized(); err != nil {
return ApplyResult{}, err
}
result := ApplyResult{
NodeID: plan.NodeID,
Script: RenderTable(plan),
Hashes: PlanHashes(plan),
}
if err := m.runner.ApplyScript(ctx, cfg, result.Script); err != nil {
return ApplyResult{}, err
}
return result, nil
}
func (m *Manager) Clear(ctx context.Context, cfg SSHConfig) error {
if err := m.ensureInitialized(); err != nil {
return err
}
script := RenderTable(NodePlan{})
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")
}
return nil
}
@@ -0,0 +1,162 @@
package nftables
import (
"context"
"errors"
"strings"
"testing"
)
type fakeRunner struct {
scripts []string
err error
testErr error
listJSON []byte
listJSONErr error
}
func (f *fakeRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
f.scripts = append(f.scripts, script)
return f.err
}
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"}}},
}
result, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, plan)
if err != nil {
t.Fatalf("Reconcile: %v", err)
}
if len(runner.scripts) != 1 {
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
}
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] == "" {
t.Fatalf("unexpected result: %+v", result)
}
}
func TestManagerReconcileReturnsRunnerError(t *testing.T) {
runner := &fakeRunner{err: errors.New("ssh failed")}
manager := NewManager(runner)
_, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, NodePlan{NodeID: 7})
if !errors.Is(err, runner.err) {
t.Fatalf("expected original runner error, got %v", err)
}
}
func TestManagerClearAppliesEmptyTable(t *testing.T) {
runner := &fakeRunner{}
manager := NewManager(runner)
if err := manager.Clear(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}); err != nil {
t.Fatalf("Clear: %v", err)
}
if len(runner.scripts) != 1 {
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
}
if strings.Contains(runner.scripts[0], "masquerade comment") {
t.Fatalf("empty table should not include masquerade:\n%s", runner.scripts[0])
}
}
func TestManagerTestPassesThroughRunnerError(t *testing.T) {
runner := &fakeRunner{testErr: errors.New("probe failed")}
manager := NewManager(runner)
err := manager.Test(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"})
if !errors.Is(err, runner.testErr) {
t.Fatalf("expected original runner error, got %v", err)
}
}
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}
expected := errors.New("nftables manager not initialized")
var nilManager *Manager
if err := nilManager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from nil manager Test, got %v", err)
}
if _, err := nilManager.Reconcile(context.Background(), cfg, plan); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from nil manager Reconcile, got %v", err)
}
if err := nilManager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
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)
}
if _, err := manager.Reconcile(context.Background(), cfg, plan); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from Reconcile, got %v", err)
}
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)
}
}
@@ -0,0 +1,51 @@
package nftables
import (
"fmt"
"net"
"net/url"
"strconv"
"strings"
)
func ParseSingleTarget(raw string) (Target, error) {
value := strings.TrimSpace(raw)
if value == "" {
return Target{}, fmt.Errorf("目标地址不能为空")
}
if strings.Contains(value, ",") || strings.Contains(value, "\n") {
return Target{}, fmt.Errorf("nftables 纯转发第一阶段仅支持单目标")
}
if hasScheme(value) {
return Target{}, fmt.Errorf("目标地址必须是 host:port,不能包含 URL scheme")
}
host, portText, err := net.SplitHostPort(value)
if err != nil {
return Target{}, fmt.Errorf("目标地址必须是 host:port")
}
host = strings.TrimSpace(strings.Trim(host, "[]"))
if host == "" {
return Target{}, fmt.Errorf("目标主机不能为空")
}
port, err := strconv.Atoi(portText)
if err != nil || port < 1 || port > 65535 {
return Target{}, fmt.Errorf("目标端口必须在 1-65535 之间")
}
return Target{Host: host, Port: port}, nil
}
func hasScheme(value string) bool {
parsed, err := url.Parse(value)
if err != nil || parsed.Scheme == "" {
return false
}
if strings.Contains(value, "://") {
return true
}
colon := strings.IndexByte(value, ':')
if colon <= 0 || strings.Contains(parsed.Scheme, ".") {
return false
}
suffix := value[colon+1:]
return strings.IndexByte(suffix, ':') == -1
}
@@ -0,0 +1,46 @@
package nftables
import "testing"
func TestParseSingleTargetAcceptsHostPortAndIPv6(t *testing.T) {
tests := []struct {
name string
raw string
host string
port int
}{
{name: "hostname", raw: "example.com:443", host: "example.com", port: 443},
{name: "ipv4", raw: "198.51.100.20:8443", host: "198.51.100.20", port: 8443},
{name: "ipv6", raw: "[2001:db8::1]:443", host: "2001:db8::1", port: 443},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
target, err := ParseSingleTarget(tt.raw)
if err != nil {
t.Fatalf("ParseSingleTarget: %v", err)
}
if target.Host != tt.host || target.Port != tt.port {
t.Fatalf("expected %s/%d, got %+v", tt.host, tt.port, target)
}
})
}
}
func TestParseSingleTargetRejectsUnsupportedValues(t *testing.T) {
for _, raw := range []string{
"",
"example.com",
"example.com:0",
"example.com:65536",
"a:1,b:2",
"http://example.com:443",
"https:443",
"mailto:443",
} {
t.Run(raw, func(t *testing.T) {
if _, err := ParseSingleTarget(raw); err == nil {
t.Fatalf("expected error for %q", raw)
}
})
}
}
@@ -0,0 +1,146 @@
package nftables
import (
"crypto/sha256"
"encoding/hex"
"fmt"
"net"
"sort"
"strings"
)
func RenderTable(plan NodePlan) string {
var b strings.Builder
b.WriteString("table inet flvx {\n")
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 counter dnat %sto %s comment %q\n",
protocol,
rule.InPort,
dnatFamily,
formatDNATTarget(rule.TargetHost, rule.TargetPort),
counterComment(rule.ForwardID, CounterDirectionDNAT, protocol),
))
}
}
b.WriteString(" }\n\n")
b.WriteString(" chain postrouting {\n")
b.WriteString(" type nat hook postrouting priority srcnat; policy accept;\n")
if len(plan.Rules) > 0 {
b.WriteString(" masquerade comment \"flvx masquerade\"\n")
}
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(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
rule.InPort,
family,
targetHost,
protocol,
rule.TargetPort,
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
))
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
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",
rule.ForwardID,
rule.InPort,
strings.TrimSpace(rule.TargetHost),
rule.TargetPort,
strings.Join(protocols, ","),
)))
return hex.EncodeToString(sum[:])
}
func PlanHashes(plan NodePlan) map[int64]string {
hashes := make(map[int64]string, len(plan.Rules))
for _, rule := range plan.Rules {
hashes[rule.ForwardID] = RuleHash(rule)
}
return hashes
}
func sortedRules(rules []Rule) []Rule {
out := append([]Rule(nil), rules...)
sort.SliceStable(out, func(i, j int) bool {
if out[i].InPort == out[j].InPort {
return out[i].ForwardID < out[j].ForwardID
}
return out[i].InPort < out[j].InPort
})
return out
}
func normalizedProtocols(protocols []string) []string {
seen := map[string]struct{}{}
out := make([]string, 0, 2)
for _, protocol := range protocols {
p := strings.ToLower(strings.TrimSpace(protocol))
if p != "tcp" && p != "udp" {
continue
}
if _, ok := seen[p]; ok {
continue
}
seen[p] = struct{}{}
out = append(out, p)
}
if len(out) == 0 {
return []string{"tcp", "udp"}
}
sort.Strings(out)
return out
}
func formatDNATTarget(host string, port int) string {
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
if ip := net.ParseIP(trimmed); ip != nil && ip.To4() == nil {
return fmt.Sprintf("[%s]:%d", trimmed, port)
}
return fmt.Sprintf("%s:%d", trimmed, port)
}
func nftAddressFamily(host string) string {
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
ip := net.ParseIP(trimmed)
if ip == nil {
return ""
}
if ip.To4() == nil {
return "ip6"
}
return "ip"
}
@@ -0,0 +1,176 @@
package nftables
import (
"strings"
"testing"
)
func TestRenderTableIncludesDNATAndMasquerade(t *testing.T) {
script := RenderTable(NodePlan{
NodeID: 10,
Rules: []Rule{
{
ForwardID: 42,
InPort: 24000,
TargetHost: "198.51.100.20",
TargetPort: 443,
Protocols: []string{"tcp", "udp"},
},
},
})
expectedParts := []string{
"table inet flvx",
"type nat hook prerouting priority dstnat; policy accept;",
"type nat hook postrouting priority srcnat; policy accept;",
"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 {
if !strings.Contains(script, part) {
t.Fatalf("script missing %q:\n%s", part, script)
}
}
}
func TestRenderTableBracketsIPv6Target(t *testing.T) {
script := RenderTable(NodePlan{
NodeID: 10,
Rules: []Rule{
{ForwardID: 42, InPort: 24000, TargetHost: "2001:db8::1", TargetPort: 443, Protocols: []string{"tcp"}},
},
})
if !strings.Contains(script, "dnat ip6 to [2001:db8::1]:443") {
t.Fatalf("expected bracketed IPv6 dnat target, got:\n%s", script)
}
}
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"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target 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"`,
`ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target 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{
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target 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) {
t.Fatalf("expected stable rule hash")
}
if RuleHash(rule) == RuleHash(Rule{ForwardID: 42, InPort: 24001, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}) {
t.Fatalf("expected hash to change when port changes")
}
}
func TestRuleHashIgnoresBindIPWhenRenderingDoesNotUseIt(t *testing.T) {
base := Rule{
ForwardID: 42,
InPort: 24000,
TargetHost: "198.51.100.20",
TargetPort: 443,
Protocols: []string{"tcp", "udp"},
}
withBind := base
withBind.BindIP = "192.0.2.10"
if RuleHash(base) != RuleHash(withBind) {
t.Fatalf("expected bind IP to be ignored by hash when it is not rendered")
}
}
@@ -0,0 +1,206 @@
package nftables
import (
"bytes"
"context"
"fmt"
"net"
"strings"
"time"
"golang.org/x/crypto/ssh"
)
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 {
Timeout time.Duration
}
func NewSSHRunner() *SSHRunner {
return &SSHRunner{Timeout: 15 * time.Second}
}
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")
}
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\""
return r.run(ctx, cfg, command)
}
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
}
runCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
clientConfig, err := buildSSHClientConfig(cfg)
if err != nil {
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 nil, fmt.Errorf("SSH 连接失败: %w", err)
}
defer conn.Close()
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
if err != nil {
return nil, fmt.Errorf("SSH 认证失败: %w", err)
}
client := ssh.NewClient(sshConn, chans, reqs)
defer client.Close()
session, err := client.NewSession()
if err != nil {
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)
go func() {
done <- session.Run(command)
}()
select {
case <-runCtx.Done():
_ = session.Close()
return nil, fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
case err := <-done:
if err != nil {
message := strings.TrimSpace(stderr.String())
if message != "" {
return nil, fmt.Errorf("远程执行失败: %s: %w", message, err)
}
return nil, fmt.Errorf("远程执行失败: %w", err)
}
return stdout.Bytes(), nil
}
}
func buildSSHClientConfig(cfg SSHConfig) (*ssh.ClientConfig, error) {
if strings.TrimSpace(cfg.Host) == "" {
return nil, fmt.Errorf("SSH 主机不能为空")
}
if strings.TrimSpace(cfg.Username) == "" {
return nil, fmt.Errorf("SSH 用户名不能为空")
}
auth, err := authMethods(cfg)
if err != nil {
return nil, err
}
if len(auth) == 0 {
return nil, fmt.Errorf("SSH 认证方式不能为空")
}
return &ssh.ClientConfig{
User: strings.TrimSpace(cfg.Username),
Auth: auth,
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
Timeout: 15 * time.Second,
}, nil
}
func authMethods(cfg SSHConfig) ([]ssh.AuthMethod, error) {
switch strings.ToLower(strings.TrimSpace(cfg.AuthType)) {
case "":
if strings.TrimSpace(cfg.PrivateKey) == "" {
return nil, fmt.Errorf("SSH 私钥不能为空")
}
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
if err != nil {
return nil, err
}
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
case "password":
if cfg.Password == "" {
return nil, fmt.Errorf("SSH 密码不能为空")
}
return []ssh.AuthMethod{ssh.Password(cfg.Password)}, nil
case "private_key":
if strings.TrimSpace(cfg.PrivateKey) == "" {
return nil, fmt.Errorf("SSH 私钥不能为空")
}
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
if err != nil {
return nil, err
}
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
default:
return nil, fmt.Errorf("不支持的 SSH 认证方式: %s", cfg.AuthType)
}
}
func parsePrivateKey(privateKey, passphrase string) (ssh.Signer, error) {
if passphrase != "" {
signer, err := ssh.ParsePrivateKeyWithPassphrase([]byte(privateKey), []byte(passphrase))
if err != nil {
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
}
return signer, nil
}
signer, err := ssh.ParsePrivateKey([]byte(privateKey))
if err != nil {
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
}
return signer, nil
}
func nftCommand(cfg SSHConfig, command string) string {
return "sh -lc " + sshQuote(command)
}
func nftBinary(cfg SSHConfig) string {
if strings.EqualFold(strings.TrimSpace(cfg.SudoMode), "sudo") {
return "sudo -n nft"
}
return "nft"
}
func sshQuote(value string) string {
return "'" + strings.ReplaceAll(value, "'", "'\"'\"'") + "'"
}
func normalizedSSHPort(port int) int {
if port <= 0 {
return 22
}
return port
}
@@ -0,0 +1,42 @@
package nftables
import (
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"strings"
"testing"
)
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
privateKey := mustGeneratePrivateKey(t)
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
if err != nil {
t.Fatalf("authMethods: %v", err)
}
if len(methods) != 1 {
t.Fatalf("expected 1 auth method, got %d", len(methods))
}
}
func TestAuthMethodsDefaultPrivateKeyRequiresKey(t *testing.T) {
_, err := authMethods(SSHConfig{})
if err == nil || !strings.Contains(err.Error(), "SSH 私钥不能为空") {
t.Fatalf("expected private key required error, got %v", err)
}
}
func mustGeneratePrivateKey(t *testing.T) string {
t.Helper()
key, err := rsa.GenerateKey(rand.Reader, 1024)
if err != nil {
t.Fatalf("GenerateKey: %v", err)
}
block := &pem.Block{
Type: "RSA PRIVATE KEY",
Bytes: x509.MarshalPKCS1PrivateKey(key),
}
return string(pem.EncodeToMemory(block))
}
@@ -0,0 +1,50 @@
package nftables
const (
ModeAgent = "agent"
ModeNftables = "nftables"
StatusPending = "pending"
StatusApplied = "applied"
StatusError = "error"
CounterDirectionDNAT = "dnat"
CounterDirectionToTarget = "to-target"
CounterDirectionFromTarget = "from-target"
)
type Target struct {
Host string
Port int
}
type Rule struct {
ForwardID int64
InPort int
BindIP string
TargetHost string
TargetPort int
Protocols []string
}
type NodePlan struct {
NodeID int64
Rules []Rule
}
type SSHConfig struct {
Host string
Port int
Username string
AuthType string
Password string
PrivateKey string
Passphrase string
SudoMode string
}
type ApplyResult struct {
NodeID int64
Script string
Hashes map[int64]string
}
+108 -49
View File
@@ -31,24 +31,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" }
@@ -87,6 +89,7 @@ type Node struct {
UDPListenAddr string `gorm:"column:udp_listen_addr;type:varchar(100);not null;default:'[::]'"`
Inx int `gorm:"not null;default:0"`
IsRemote int `gorm:"column:is_remote;default:0"`
ForwardMode string `gorm:"column:forward_mode;type:varchar(20);not null;default:'agent'"`
RemoteURL sql.NullString `gorm:"column:remote_url;type:text"`
RemoteToken sql.NullString `gorm:"column:remote_token;type:text"`
RemoteConfig sql.NullString `gorm:"column:remote_config;type:text"`
@@ -95,6 +98,57 @@ type Node struct {
func (Node) TableName() string { return "node" }
type NodeSSHConfig struct {
ID int64 `gorm:"primaryKey;autoIncrement"`
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex"`
Host string `gorm:"type:varchar(255);not null"`
Port int `gorm:"not null;default:22"`
Username string `gorm:"type:varchar(100);not null"`
AuthType string `gorm:"column:auth_type;type:varchar(20);not null"`
Password sql.NullString `gorm:"type:text"`
PrivateKey sql.NullString `gorm:"column:private_key;type:text"`
Passphrase sql.NullString `gorm:"type:text"`
SudoMode string `gorm:"column:sudo_mode;type:varchar(20);not null;default:'none'"`
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime int64 `gorm:"column:updated_time;not null"`
}
func (NodeSSHConfig) TableName() string { return "node_ssh_config" }
type NftRuleBinding struct {
ID int64 `gorm:"primaryKey;autoIncrement"`
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_rule_binding_forward_node;index"`
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_rule_binding_forward_node;index"`
InPort int `gorm:"column:in_port;not null"`
Protocols string `gorm:"type:varchar(20);not null;default:'tcp,udp'"`
TargetAddr string `gorm:"column:target_addr;type:text;not null"`
BindIP string `gorm:"column:bind_ip;type:text;not null;default:''"`
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
Status string `gorm:"type:varchar(20);not null;default:'pending'"`
LastError string `gorm:"column:last_error;type:text;not null;default:''"`
AppliedTime int64 `gorm:"column:applied_time;not null;default:0"`
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime int64 `gorm:"column:updated_time;not null"`
}
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"`
@@ -435,24 +489,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 {
@@ -541,19 +597,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.
@@ -602,6 +660,7 @@ type NodeRecord struct {
UDPListenAddr string
InterfaceName string
IsRemote int
ForwardMode string
RemoteURL string
RemoteToken string
RemoteConfig string
+140 -73
View File
@@ -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 ────────────────────────────────────────────────────
@@ -189,7 +200,6 @@ func Open(path string) (*Repository, error) {
_ = sqlDB.Close()
return nil, fmt.Errorf("prepare sqlite legacy schema: %w", err)
}
if err := autoMigrateAll(db); err != nil {
_ = sqlDB.Close()
return nil, fmt.Errorf("auto migrate: %w", err)
@@ -269,12 +279,21 @@ func (r *Repository) Close() error {
}
func autoMigrateAll(db *gorm.DB) error {
if db.Dialector.Name() == "sqlite" {
if err := prepareSQLiteNftablesColumns(db); err != nil {
return err
}
}
models := []interface{}{
&model.User{},
&model.UserQuota{},
&model.Forward{},
&model.ForwardPort{},
&model.Node{},
&model.NodeSSHConfig{},
&model.NftRuleBinding{},
&model.NftCounterState{},
&model.SpeedLimit{},
&model.StatisticsFlow{},
&model.Tunnel{},
@@ -405,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
}
@@ -418,6 +437,21 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
return nil
}
func prepareSQLiteNftablesColumns(db *gorm.DB) error {
if db == nil || db.Dialector.Name() != "sqlite" {
return nil
}
if !db.Migrator().HasTable(&model.Node{}) {
return nil
}
if !db.Migrator().HasColumn(&model.Node{}, "forward_mode") {
if err := db.Exec("ALTER TABLE node ADD COLUMN forward_mode varchar(20) NOT NULL DEFAULT 'agent'").Error; err != nil {
return err
}
}
return nil
}
func seedData(db *gorm.DB) {
var adminCount int64
if err := db.Model(&model.User{}).Where("id = ?", 1).Count(&adminCount).Error; err == nil && adminCount == 0 {
@@ -793,6 +827,28 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
if err := r.db.Order("inx ASC, id ASC").Find(&nodes).Error; err != nil {
return nil, err
}
nodeIDs := make([]int64, 0, len(nodes))
for _, n := range nodes {
if defaultNodeForwardMode(n.ForwardMode) == "nftables" {
nodeIDs = append(nodeIDs, n.ID)
}
}
sshConfigByNodeID := make(map[int64]map[string]interface{}, len(nodeIDs))
if len(nodeIDs) > 0 {
var configs []model.NodeSSHConfig
if err := r.db.Where("node_id IN ?", nodeIDs).Find(&configs).Error; err != nil {
return nil, err
}
for _, cfg := range configs {
sshConfigByNodeID[cfg.NodeID] = map[string]interface{}{
"host": cfg.Host,
"port": cfg.Port,
"username": cfg.Username,
"authType": cfg.AuthType,
"sudoMode": cfg.SudoMode,
}
}
}
items := make([]map[string]interface{}, 0, len(nodes))
for _, n := range nodes {
items = append(items, map[string]interface{}{
@@ -810,11 +866,13 @@ func (r *Repository) ListNodes() ([]map[string]interface{}, error) {
"version": nullableString(n.Version),
"http": n.HTTP, "tls": n.TLS, "socks": n.Socks,
"status": n.Status, "isRemote": n.IsRemote,
"forwardMode": defaultNodeForwardMode(n.ForwardMode),
"remoteUrl": nullableString(n.RemoteURL),
"remoteToken": nullableString(n.RemoteToken),
"remoteConfig": nullableString(n.RemoteConfig),
"expiryReminderDismissed": n.ExpiryReminderDismissed,
"interfaceName": nullableString(n.InterfaceName),
"sshConfig": sshConfigByNodeID[n.ID],
})
}
return items, nil
@@ -886,31 +944,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").
@@ -921,6 +981,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
@@ -933,9 +994,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
@@ -2137,8 +2200,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
@@ -2572,29 +2637,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 {
@@ -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 {
@@ -233,7 +236,8 @@ func nodeRecordFromModel(n *model.Node) *model.NodeRecord {
Status: n.Status,
PortRange: n.Port,
TCPListenAddr: n.TCPListenAddr, UDPListenAddr: n.UDPListenAddr,
IsRemote: n.IsRemote,
IsRemote: n.IsRemote,
ForwardMode: defaultNodeForwardMode(n.ForwardMode),
}
if n.ServerIPV4.Valid {
rec.ServerIPv4 = strings.TrimSpace(n.ServerIPV4.String)
@@ -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])
}
}
@@ -197,6 +197,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 ""
@@ -220,7 +233,7 @@ func (r *Repository) GetUserDefaultsForTunnel(userID int64) (flow int64, num int
return user.Flow, user.Num, user.ExpTime, user.FlowResetTime, nil
}
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}) error {
func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serverIPV6, port, interfaceName, version, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, now int64, status int, tcpAddr, udpAddr string, inx, isRemote int, remoteURL, remoteToken, remoteConfig, extraIPs interface{}, forwardMode string) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
@@ -247,6 +260,7 @@ func (r *Repository) CreateNode(name, secret, serverIP string, serverIPV4, serve
UDPListenAddr: udpAddr,
Inx: inx,
IsRemote: isRemote,
ForwardMode: defaultNodeForwardMode(forwardMode),
RemoteURL: nullStringFromInterface(remoteURL),
RemoteToken: nullStringFromInterface(remoteToken),
RemoteConfig: nullStringFromInterface(remoteConfig),
@@ -266,31 +280,44 @@ func (r *Repository) GetNodeStatusFields(nodeID int64) (status, httpFlag, tlsFla
return node.Status, node.HTTP, node.TLS, node.Socks, nil
}
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
func (r *Repository) UpdateNode(id int64, name, serverIP string, serverIPV4, serverIPV6, port, interfaceName, extraIPs, remark, expiryTime, renewalCycle interface{}, forwardMode string, httpFlag, tlsFlag, socksFlag int, tcpAddr, udpAddr string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
updates := map[string]interface{}{
"name": name,
"remark": nullStringFromInterface(remark),
"expiry_time": nullInt64FromInterface(expiryTime),
"renewal_cycle": nullStringFromInterface(renewalCycle),
"server_ip": serverIP,
"server_ip_v4": nullStringFromInterface(serverIPV4),
"server_ip_v6": nullStringFromInterface(serverIPV6),
"extra_ips": nullStringFromInterface(extraIPs),
"port": stringFromInterface(port),
"interface_name": nullStringFromInterface(interfaceName),
"http": httpFlag,
"tls": tlsFlag,
"socks": socksFlag,
"tcp_listen_addr": tcpAddr,
"udp_listen_addr": udpAddr,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
"expiry_reminder_dismissed": 0,
}
if strings.TrimSpace(forwardMode) != "" {
updates["forward_mode"] = defaultNodeForwardMode(forwardMode)
}
return r.db.Model(&model.Node{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"name": name,
"remark": nullStringFromInterface(remark),
"expiry_time": nullInt64FromInterface(expiryTime),
"renewal_cycle": nullStringFromInterface(renewalCycle),
"server_ip": serverIP,
"server_ip_v4": nullStringFromInterface(serverIPV4),
"server_ip_v6": nullStringFromInterface(serverIPV6),
"extra_ips": nullStringFromInterface(extraIPs),
"port": stringFromInterface(port),
"interface_name": nullStringFromInterface(interfaceName),
"http": httpFlag,
"tls": tlsFlag,
"socks": socksFlag,
"tcp_listen_addr": tcpAddr,
"udp_listen_addr": udpAddr,
"updated_time": sql.NullInt64{Int64: now, Valid: true},
"expiry_reminder_dismissed": 0,
}).Error
Updates(updates).Error
}
func defaultNodeForwardMode(mode string) string {
switch strings.TrimSpace(strings.ToLower(mode)) {
case "nftables":
return "nftables"
default:
return "agent"
}
}
func (r *Repository) GetNodeSecret(nodeID int64) (string, error) {
@@ -712,23 +739,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
}
@@ -755,6 +785,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
}
@@ -802,26 +835,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
}
@@ -1281,30 +1317,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)
}
}
@@ -0,0 +1,249 @@
package repo
import (
"database/sql"
"errors"
"strings"
"go-backend/internal/store/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
type NftSSHConfigInput struct {
Host string
Port int
Username string
AuthType string
Password string
PrivateKey string
Passphrase string
SudoMode string
}
type NftRuleBindingInput struct {
ForwardID int64
NodeID int64
InPort int
Protocols string
TargetAddr string
BindIP string
RuleHash string
Status string
LastError string
}
func (r *Repository) UpsertNodeSSHConfig(nodeID int64, cfg NftSSHConfigInput, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if nodeID <= 0 {
return errors.New("node id is required")
}
port := cfg.Port
if port <= 0 {
port = 22
}
authType := strings.TrimSpace(strings.ToLower(cfg.AuthType))
if authType == "" {
authType = "private_key"
}
sudoMode := strings.TrimSpace(strings.ToLower(cfg.SudoMode))
if sudoMode == "" {
sudoMode = "none"
}
row := model.NodeSSHConfig{
NodeID: nodeID,
Host: strings.TrimSpace(cfg.Host),
Port: port,
Username: strings.TrimSpace(cfg.Username),
AuthType: authType,
Password: nullStringFromInterface(cfg.Password),
PrivateKey: nullStringFromInterface(cfg.PrivateKey),
Passphrase: nullStringFromInterface(cfg.Passphrase),
SudoMode: sudoMode,
CreatedTime: now,
UpdatedTime: now,
}
return r.db.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "node_id"}},
DoUpdates: clause.Assignments(map[string]interface{}{
"host": row.Host,
"port": row.Port,
"username": row.Username,
"auth_type": row.AuthType,
"password": row.Password,
"private_key": row.PrivateKey,
"passphrase": row.Passphrase,
"sudo_mode": row.SudoMode,
"updated_time": row.UpdatedTime,
}),
}).Create(&row).Error
}
func (r *Repository) GetNodeSSHConfig(nodeID int64) (*model.NodeSSHConfig, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
return r.GetNodeSSHConfigTx(r.db, nodeID)
}
func (r *Repository) GetNodeSSHConfigTx(tx *gorm.DB, nodeID int64) (*model.NodeSSHConfig, error) {
if tx == nil {
return nil, errors.New("database unavailable")
}
var cfg model.NodeSSHConfig
if err := tx.Where("node_id = ?", nodeID).First(&cfg).Error; err != nil {
return nil, normalizeNotFoundErr(err)
}
return &cfg, nil
}
func (r *Repository) DeleteNodeSSHConfig(nodeID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Where("node_id = ?", nodeID).Delete(&model.NodeSSHConfig{}).Error
}
func (r *Repository) UpsertNftRuleBinding(input NftRuleBindingInput, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
row := model.NftRuleBinding{
ForwardID: input.ForwardID,
NodeID: input.NodeID,
InPort: input.InPort,
Protocols: defaultString(strings.TrimSpace(input.Protocols), "tcp,udp"),
TargetAddr: strings.TrimSpace(input.TargetAddr),
BindIP: strings.TrimSpace(input.BindIP),
RuleHash: strings.TrimSpace(input.RuleHash),
Status: defaultString(strings.TrimSpace(input.Status), "pending"),
LastError: strings.TrimSpace(input.LastError),
AppliedTime: now,
CreatedTime: now,
UpdatedTime: now,
}
return r.db.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "forward_id"}, {Name: "node_id"}},
DoUpdates: clause.Assignments(map[string]interface{}{
"in_port": row.InPort,
"protocols": row.Protocols,
"target_addr": row.TargetAddr,
"bind_ip": row.BindIP,
"rule_hash": row.RuleHash,
"status": row.Status,
"last_error": row.LastError,
"applied_time": row.AppliedTime,
"updated_time": row.UpdatedTime,
}),
}).Create(&row).Error
}
func (r *Repository) MarkNftRuleBindingError(forwardID, nodeID int64, message string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Model(&model.NftRuleBinding{}).
Where("forward_id = ? AND node_id = ?", forwardID, nodeID).
Updates(map[string]interface{}{
"status": "error",
"last_error": strings.TrimSpace(message),
"updated_time": now,
}).Error
}
func (r *Repository) ListNftRuleBindingsByNode(nodeID int64) ([]model.NftRuleBinding, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var rows []model.NftRuleBinding
err := r.db.Where("node_id = ?", nodeID).Order("forward_id ASC").Find(&rows).Error
return rows, err
}
func (r *Repository) DeleteNftRuleBindingsByForward(forwardID int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftRuleBinding{}).Error
}
func (r *Repository) GetNodeForwardMode(nodeID int64) (string, error) {
if r == nil || r.db == nil {
return "", errors.New("repository not initialized")
}
return r.GetNodeForwardModeTx(r.db, nodeID)
}
func (r *Repository) GetNodeForwardModeTx(tx *gorm.DB, nodeID int64) (string, error) {
if tx == nil {
return "", errors.New("database unavailable")
}
var row struct {
ForwardMode sql.NullString `gorm:"column:forward_mode"`
}
err := tx.Model(&model.Node{}).Select("forward_mode").Where("id = ?", nodeID).First(&row).Error
if err != nil {
return "", normalizeNotFoundErr(err)
}
return defaultNodeForwardMode(row.ForwardMode.String), nil
}
func (r *Repository) ListActiveForwardsByNode(nodeID int64) ([]model.ForwardRecord, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var forwards []model.Forward
err := r.db.Model(&model.Forward{}).
Joins("JOIN forward_port ON forward_port.forward_id = forward.id").
Where("forward_port.node_id = ? AND forward.status = 1", nodeID).
Order("forward.id ASC").
Distinct("forward.*").
Find(&forwards).Error
if err != nil {
return nil, err
}
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,
ProxyProtocolReceive: proxyProtocolReceive,
ProxyProtocolSend: proxyProtocolSend,
})
}
for i := range rows {
if strings.TrimSpace(rows[i].Strategy) == "" {
rows[i].Strategy = "fifo"
}
}
return rows, nil
}
func defaultString(value, fallback string) string {
if strings.TrimSpace(value) == "" {
return fallback
}
return value
}
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
if send == 0 && legacy > 0 {
send = legacy
}
return receive, send
}
@@ -0,0 +1,323 @@
package repo
import (
"path/filepath"
"strings"
"testing"
"time"
)
func TestNftablesNodeModeSSHConfigAndBindingPersistence(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.CreateNode(
"nft-node",
"secret",
"203.0.113.10",
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)
}
if len(nodes) != 1 {
t.Fatalf("expected 1 node, got %d", len(nodes))
}
nodeID := nodes[0]["id"].(int64)
if got := nodes[0]["forwardMode"]; got != "nftables" {
t.Fatalf("expected forwardMode nftables, got %#v", got)
}
cfg := NftSSHConfigInput{
Host: "203.0.113.10",
Port: 22,
Username: "root",
AuthType: "private_key",
PrivateKey: "encrypted-private-key",
SudoMode: "none",
}
if err := r.UpsertNodeSSHConfig(nodeID, cfg, now); err != nil {
t.Fatalf("UpsertNodeSSHConfig: %v", err)
}
loaded, err := r.GetNodeSSHConfig(nodeID)
if err != nil {
t.Fatalf("GetNodeSSHConfig: %v", err)
}
if loaded.Host != cfg.Host || loaded.Port != cfg.Port || loaded.Username != cfg.Username || loaded.AuthType != cfg.AuthType {
t.Fatalf("unexpected ssh config: %+v", loaded)
}
binding := NftRuleBindingInput{
ForwardID: 42,
NodeID: nodeID,
InPort: 24000,
Protocols: "tcp,udp",
TargetAddr: "198.51.100.20:443",
BindIP: "",
RuleHash: "hash-a",
Status: "applied",
LastError: "",
}
if err := r.UpsertNftRuleBinding(binding, now); err != nil {
t.Fatalf("UpsertNftRuleBinding: %v", err)
}
bindings, err := r.ListNftRuleBindingsByNode(nodeID)
if err != nil {
t.Fatalf("ListNftRuleBindingsByNode: %v", err)
}
if len(bindings) != 1 {
t.Fatalf("expected 1 binding, got %d", len(bindings))
}
if bindings[0].ForwardID != 42 || bindings[0].RuleHash != "hash-a" || bindings[0].Status != "applied" {
t.Fatalf("unexpected binding: %+v", bindings[0])
}
if err := r.MarkNftRuleBindingError(42, nodeID, "nft failed", now+1); err != nil {
t.Fatalf("MarkNftRuleBindingError: %v", err)
}
bindings, err = r.ListNftRuleBindingsByNode(nodeID)
if err != nil {
t.Fatalf("ListNftRuleBindingsByNode after error: %v", err)
}
if bindings[0].Status != "error" || !strings.Contains(bindings[0].LastError, "nft failed") {
t.Fatalf("expected error binding, got %+v", bindings[0])
}
if err := r.DeleteNftRuleBindingsByForward(42); err != nil {
t.Fatalf("DeleteNftRuleBindingsByForward: %v", err)
}
bindings, err = r.ListNftRuleBindingsByNode(nodeID)
if err != nil {
t.Fatalf("ListNftRuleBindingsByNode after delete: %v", err)
}
if len(bindings) != 0 {
t.Fatalf("expected no bindings after delete, got %+v", bindings)
}
}
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 {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
if err := r.CreateNode(
"nft-node",
"secret",
"203.0.113.11",
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)
}
if len(nodes) != 1 {
t.Fatalf("expected 1 node, got %d", len(nodes))
}
nodeID := nodes[0]["id"].(int64)
if err := r.UpdateNode(
nodeID,
"nft-node-updated",
"203.0.113.11",
nil,
nil,
"10000-20000",
nil,
nil,
nil,
nil,
nil,
"",
0,
0,
0,
"[::]",
"[::]",
now+1,
); err != nil {
t.Fatalf("UpdateNode: %v", err)
}
gotNode, err := r.GetNodeRecord(nodeID)
if err != nil {
t.Fatalf("GetNodeRecord: %v", err)
}
if gotNode == nil {
t.Fatal("expected node record, got nil")
}
if gotNode.ForwardMode != "nftables" {
t.Fatalf("expected mapped forward mode nftables, got %q", gotNode.ForwardMode)
}
nodes, err = r.ListNodes()
if err != nil {
t.Fatalf("ListNodes after update: %v", err)
}
if got := nodes[0]["forwardMode"]; got != "nftables" {
t.Fatalf("expected persisted forwardMode nftables after update, got %#v", got)
}
}
@@ -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")
+31 -1
View File
@@ -151,6 +151,7 @@ const (
initialBackoff = 2 * time.Second // 重连初始退避
maxBackoff = 2 * time.Minute // 重连最大退避
defaultMetricReportInterval = 5 * time.Second
maxConcurrentTCPPings = 8
)
type WebSocketReporter struct {
@@ -172,6 +173,7 @@ type WebSocketReporter struct {
connecting bool // 正在连接状态
connMutex sync.Mutex // 连接状态锁
aesCrypto *crypto.AESCrypto // AES加密器
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
}
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
@@ -201,6 +203,29 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
connected: false,
connecting: false,
aesCrypto: aesCrypto,
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
}
}
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:
}
}
@@ -840,9 +865,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{
+247 -40
View File
@@ -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,11 +51,14 @@ 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:-}"
@@ -256,6 +286,199 @@ 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
if command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
SERVICE_MANAGER="systemd"
return 0
fi
if 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
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=""
@@ -341,19 +564,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 +617,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 +651,7 @@ update_flux_agent() {
# 检查并安装 tcpkill
check_and_install_tcpkill
ensure_service_manager || return 1
# 先下载新版本
echo "⬇️ 下载最新版本..."
@@ -455,9 +664,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 +678,7 @@ update_flux_agent() {
# 重启服务
echo "🔄 重启服务..."
systemctl start flux_agent
start_flux_agent_service
echo "✅ 更新完成,服务已重新启动。"
}
@@ -477,6 +686,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 +695,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 +713,6 @@ uninstall_flux_agent() {
echo "🧹 删除安装目录: $INSTALL_DIR"
fi
# 重载 systemd
systemctl daemon-reload
echo "✅ 卸载完成"
}
+120 -1
View File
@@ -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"
@@ -274,6 +279,117 @@ EOF
assert_equals "$expected" "$actual" "install_flux_agent should JSON-escape config values"
)
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() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
@@ -504,6 +620,9 @@ 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_install_script_accepts_proxy_url_env_without_prompt
@@ -514,4 +633,4 @@ 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"
+1 -20
View File
@@ -1496,28 +1496,24 @@ packages:
engines: {node: ^20.19.0 || >=22.12.0}
cpu: [arm64]
os: [linux]
libc: [glibc]
'@rolldown/binding-linux-arm64-musl@1.0.0-beta.53':
resolution: {integrity: sha512-bGe5EBB8FVjHBR1mOLOPEFg1Lp3//7geqWkU5NIhxe+yH0W8FVrQ6WRYOap4SUTKdklD/dC4qPLREkMMQ855FA==}
engines: {node: ^20.19.0 || >=22.12.0}
cpu: [arm64]
os: [linux]
libc: [musl]
'@rolldown/binding-linux-x64-gnu@1.0.0-beta.53':
resolution: {integrity: sha512-qL+63WKVQs1CMvFedlPt0U9PiEKJOAL/bsHMKUDS6Vp2Q+YAv/QLPu8rcvkfIMvQ0FPU2WL0aX4eWwF6e/GAnA==}
engines: {node: ^20.19.0 || >=22.12.0}
cpu: [x64]
os: [linux]
libc: [glibc]
'@rolldown/binding-linux-x64-musl@1.0.0-beta.53':
resolution: {integrity: sha512-VGl9JIGjoJh3H8Mb+7xnVqODajBmrdOOb9lxWXdcmxyI+zjB2sux69br0hZJDTyLJfvBoYm439zPACYbCjGRmw==}
engines: {node: ^20.19.0 || >=22.12.0}
cpu: [x64]
os: [linux]
libc: [musl]
'@rolldown/binding-openharmony-arm64@1.0.0-beta.53':
resolution: {integrity: sha512-B4iIserJXuSnNzA5xBLFUIjTfhNy7d9sq4FUMQY3GhQWGVhS2RWWzzDnkSU6MUt7/aHUrep0CdQfXUJI9D3W7A==}
@@ -1683,56 +1679,48 @@ packages:
engines: {node: '>= 10'}
cpu: [arm64]
os: [linux]
libc: [glibc]
'@tailwindcss/oxide-linux-arm64-gnu@4.2.4':
resolution: {integrity: sha512-+E4wxJ0ZGOzSH325reXTWB48l42i93kQqMvDyz5gqfRzRZ7faNhnmvlV4EPGJU3QJM/3Ab5jhJ5pCRUsKn6OQw==}
engines: {node: '>= 20'}
cpu: [arm64]
os: [linux]
libc: [glibc]
'@tailwindcss/oxide-linux-arm64-musl@4.1.11':
resolution: {integrity: sha512-m/NVRFNGlEHJrNVk3O6I9ggVuNjXHIPoD6bqay/pubtYC9QIdAMpS+cswZQPBLvVvEF6GtSNONbDkZrjWZXYNQ==}
engines: {node: '>= 10'}
cpu: [arm64]
os: [linux]
libc: [musl]
'@tailwindcss/oxide-linux-arm64-musl@4.2.4':
resolution: {integrity: sha512-bBADEGAbo4ASnppIziaQJelekCxdMaxisrk+fB7Thit72IBnALp9K6ffA2G4ruj90G9XRS2VQ6q2bCKbfFV82g==}
engines: {node: '>= 20'}
cpu: [arm64]
os: [linux]
libc: [musl]
'@tailwindcss/oxide-linux-x64-gnu@4.1.11':
resolution: {integrity: sha512-YW6sblI7xukSD2TdbbaeQVDysIm/UPJtObHJHKxDEcW2exAtY47j52f8jZXkqE1krdnkhCMGqP3dbniu1Te2Fg==}
engines: {node: '>= 10'}
cpu: [x64]
os: [linux]
libc: [glibc]
'@tailwindcss/oxide-linux-x64-gnu@4.2.4':
resolution: {integrity: sha512-7Mx25E4WTfnht0TVRTyC00j3i0M+EeFe7wguMDTlX4mRxafznw0CA8WJkFjWYH5BlgELd1kSjuU2JiPnNZbJDA==}
engines: {node: '>= 20'}
cpu: [x64]
os: [linux]
libc: [glibc]
'@tailwindcss/oxide-linux-x64-musl@4.1.11':
resolution: {integrity: sha512-e3C/RRhGunWYNC3aSF7exsQkdXzQ/M+aYuZHKnw4U7KQwTJotnWsGOIVih0s2qQzmEzOFIJ3+xt7iq67K/p56Q==}
engines: {node: '>= 10'}
cpu: [x64]
os: [linux]
libc: [musl]
'@tailwindcss/oxide-linux-x64-musl@4.2.4':
resolution: {integrity: sha512-2wwJRF7nyhOR0hhHoChc04xngV3iS+akccHTGtz965FwF0up4b2lOdo6kI1EbDaEXKgvcrFBYcYQQ/rrnWFVfA==}
engines: {node: '>= 20'}
cpu: [x64]
os: [linux]
libc: [musl]
'@tailwindcss/oxide-wasm32-wasi@4.1.11':
resolution: {integrity: sha512-Xo1+/GU0JEN/C/dvcammKHzeM6NqKovG+6921MR6oadee5XPBaKOumrJCXvopJ/Qb5TH7LX/UAywbqrP4lax0g==}
@@ -1954,6 +1942,7 @@ packages:
'@ungap/structured-clone@1.3.0':
resolution: {integrity: sha512-WmoN8qaIAo7WTYWbAZuG8PYEhn5fkz7dZrqTBZ7dtt//lL2Gwms1IcnQ5yHqjDfX8Ft5j4YzDM23f87zBfDe9g==}
deprecated: Potential CWE-502 - Update to 1.3.1 or higher
'@vitejs/plugin-react@5.2.0':
resolution: {integrity: sha512-YmKkfhOAi3wsB1PhJq5Scj3GXMn3WvtQ/JC0xoopuHoXSdmtdStOpFrYaT1kie2YgFBcIe64ROzMYRjCrYOdYw==}
@@ -3077,56 +3066,48 @@ packages:
engines: {node: '>= 12.0.0'}
cpu: [arm64]
os: [linux]
libc: [glibc]
lightningcss-linux-arm64-gnu@1.32.0:
resolution: {integrity: sha512-0nnMyoyOLRJXfbMOilaSRcLH3Jw5z9HDNGfT/gwCPgaDjnx0i8w7vBzFLFR1f6CMLKF8gVbebmkUN3fa/kQJpQ==}
engines: {node: '>= 12.0.0'}
cpu: [arm64]
os: [linux]
libc: [glibc]
lightningcss-linux-arm64-musl@1.30.1:
resolution: {integrity: sha512-jmUQVx4331m6LIX+0wUhBbmMX7TCfjF5FoOH6SD1CttzuYlGNVpA7QnrmLxrsub43ClTINfGSYyHe2HWeLl5CQ==}
engines: {node: '>= 12.0.0'}
cpu: [arm64]
os: [linux]
libc: [musl]
lightningcss-linux-arm64-musl@1.32.0:
resolution: {integrity: sha512-UpQkoenr4UJEzgVIYpI80lDFvRmPVg6oqboNHfoH4CQIfNA+HOrZ7Mo7KZP02dC6LjghPQJeBsvXhJod/wnIBg==}
engines: {node: '>= 12.0.0'}
cpu: [arm64]
os: [linux]
libc: [musl]
lightningcss-linux-x64-gnu@1.30.1:
resolution: {integrity: sha512-piWx3z4wN8J8z3+O5kO74+yr6ze/dKmPnI7vLqfSqI8bccaTGY5xiSGVIJBDd5K5BHlvVLpUB3S2YCfelyJ1bw==}
engines: {node: '>= 12.0.0'}
cpu: [x64]
os: [linux]
libc: [glibc]
lightningcss-linux-x64-gnu@1.32.0:
resolution: {integrity: sha512-V7Qr52IhZmdKPVr+Vtw8o+WLsQJYCTd8loIfpDaMRWGUZfBOYEJeyJIkqGIDMZPwPx24pUMfwSxxI8phr/MbOA==}
engines: {node: '>= 12.0.0'}
cpu: [x64]
os: [linux]
libc: [glibc]
lightningcss-linux-x64-musl@1.30.1:
resolution: {integrity: sha512-rRomAK7eIkL+tHY0YPxbc5Dra2gXlI63HL+v1Pdi1a3sC+tJTcFrHX+E86sulgAXeI7rSzDYhPSeHHjqFhqfeQ==}
engines: {node: '>= 12.0.0'}
cpu: [x64]
os: [linux]
libc: [musl]
lightningcss-linux-x64-musl@1.32.0:
resolution: {integrity: sha512-bYcLp+Vb0awsiXg/80uCRezCYHNg1/l3mt0gzHnWV9XP1W5sKa5/TCdGWaR/zBM2PeF/HbsQv/j2URNOiVuxWg==}
engines: {node: '>= 12.0.0'}
cpu: [x64]
os: [linux]
libc: [musl]
lightningcss-win32-arm64-msvc@1.30.1:
resolution: {integrity: sha512-mSL4rqPi4iXq5YVqzSsJgMVFENoa4nGTT/GjO2c0Yl9OuQfPsIfncvLrEW6RbbB24WtZ3xP/2CCmI3tNkNV4oA==}
+2
View File
@@ -0,0 +1,2 @@
allowBuilds:
'@tailwindcss/oxide': true
+8
View File
@@ -129,6 +129,12 @@ export const getNodeReleases = (channel: ReleaseChannel = "stable") =>
Network.post<NodeReleaseApiItem[]>("/node/releases", { channel });
export const rollbackNode = (id: number) =>
Network.post("/node/rollback", { id });
export const testNodeNftables = (nodeId: number) =>
Network.post("/node/nftables/test", { nodeId });
export const reconcileNodeNftables = (nodeId: number) =>
Network.post("/node/nftables/reconcile", { nodeId });
export const clearNodeNftables = (nodeId: number) =>
Network.post("/node/nftables/clear", { nodeId });
// 隧道CRUD操作 - 全部使用POST请求
export const createTunnel = (data: TunnelMutationPayload) =>
@@ -211,6 +217,8 @@ export const pauseForwardService = (forwardId: number) =>
Network.post("/forward/pause", { id: forwardId });
export const resumeForwardService = (forwardId: number) =>
Network.post("/forward/resume", { id: forwardId });
export const resetForwardFlow = (forwardId: number) =>
Network.post("/forward/reset-flow", { id: forwardId });
// 转发诊断操作
export const diagnoseForward = (forwardId: number) =>
+49
View File
@@ -2,6 +2,8 @@ export interface NodeApiItem {
id: number;
name: string;
status: number;
forwardMode?: "agent" | "nftables";
sshConfig?: NodeSshConfigApiItem | null;
inx?: number;
remark?: string;
expiryTime?: number;
@@ -44,6 +46,7 @@ export interface TunnelApiItem {
name: string;
type: number;
status: number;
forwardMode?: "agent" | "nftables";
flow?: number;
trafficRatio?: number;
inIp?: string;
@@ -63,6 +66,7 @@ export interface ForwardApiItem {
id: number;
name: string;
status: number;
forwardMode?: "agent" | "nftables";
tunnelName?: string;
tunnelTrafficRatio?: number;
inIp?: string;
@@ -78,6 +82,8 @@ export interface ForwardApiItem {
ipSpeedLimitName?: string;
maxConn?: number;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
inx?: number;
[key: string]: unknown;
}
@@ -303,6 +309,8 @@ export interface NodeMutationPayload {
id?: number | null;
name?: string;
status?: number;
forwardMode?: "agent" | "nftables";
sshConfig?: NodeSshConfigMutationPayload | null;
inx?: number;
remark?: string;
expiryTime?: number;
@@ -320,6 +328,29 @@ export interface NodeMutationPayload {
socks?: number;
}
export interface NodeSshConfigApiItem {
host?: string;
port?: number;
username?: string;
authType?: "password" | "private_key";
password?: string;
privateKey?: string;
passphrase?: string;
sudoMode?: "none" | "sudo" | "sudo_su";
[key: string]: unknown;
}
export interface NodeSshConfigMutationPayload {
host?: string;
port?: number;
username?: string;
authType?: "password" | "private_key";
password?: string;
privateKey?: string;
passphrase?: string;
sudoMode?: "none" | "sudo" | "sudo_su";
}
export interface TunnelChainNodePayload {
nodeId: number;
protocol?: string;
@@ -334,6 +365,7 @@ export interface TunnelMutationPayload {
name?: string;
type?: number;
status?: number;
forwardMode?: "agent" | "nftables";
flow?: number;
trafficRatio?: number;
inIp?: string;
@@ -381,6 +413,7 @@ export interface ForwardMutationPayload {
id?: number;
name?: string;
status?: number;
forwardMode?: "agent" | "nftables";
tunnelId?: number | null;
inIp?: string;
inPort?: number | null;
@@ -391,6 +424,8 @@ export interface ForwardMutationPayload {
ipSpeedId?: number | null;
maxConn?: number;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
}
export interface SpeedLimitMutationPayload {
@@ -560,6 +595,20 @@ export interface TunnelQualityHopApiItem {
targetPort?: number;
}
export interface TunnelQualityCandidateHopApiItem
extends TunnelQualityHopApiItem {
fromRole: "entry" | "middle" | "exit";
toRole: "middle" | "exit" | "target";
hopIndex: number;
selected: boolean;
errorMessage?: string;
}
export interface TunnelQualityChainDetailsApiItem {
primaryPath?: TunnelQualityHopApiItem[];
candidateHops?: TunnelQualityCandidateHopApiItem[];
}
export interface TunnelQualityApiItem {
tunnelId: number;
entryToExitLatency: number;
@@ -0,0 +1,39 @@
export const TUNNEL_QUALITY_INTERVAL_CONFIG_KEY =
"monitor_tunnel_quality_interval_sec";
export const DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC = 1;
export const MIN_TUNNEL_QUALITY_INTERVAL_SEC = 1;
export const MAX_TUNNEL_QUALITY_INTERVAL_SEC = 3600;
export const parseTunnelQualityIntervalSeconds = (value: unknown): number => {
const seconds = Number(value);
return Number.isInteger(seconds) &&
seconds >= MIN_TUNNEL_QUALITY_INTERVAL_SEC &&
seconds <= MAX_TUNNEL_QUALITY_INTERVAL_SEC
? seconds
: DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC;
};
export const validateTunnelQualityInterval = (value: string): string | null => {
const normalized = value.trim();
if (!normalized) {
return "请输入探测间隔";
}
const seconds = Number(normalized);
if (!Number.isInteger(seconds)) {
return "探测间隔必须是整数";
}
if (
seconds < MIN_TUNNEL_QUALITY_INTERVAL_SEC ||
seconds > MAX_TUNNEL_QUALITY_INTERVAL_SEC
) {
return `探测间隔必须在 ${MIN_TUNNEL_QUALITY_INTERVAL_SEC} 到 ${MAX_TUNNEL_QUALITY_INTERVAL_SEC} 秒之间`;
}
return null;
};
export const tunnelQualityIntervalLabel = (seconds: number): string =>
seconds === 1 ? "每秒" : `每 ${seconds} 秒`;
+72 -4
View File
@@ -43,6 +43,14 @@ import { BackIcon, SettingsIcon } from "@/components/icons";
import { ThemeSettings } from "@/components/theme-settings";
import { isAdmin } from "@/utils/auth";
import { getCachedConfigs, configCache, updateSiteConfig } from "@/config/site";
import {
DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC,
MAX_TUNNEL_QUALITY_INTERVAL_SEC,
MIN_TUNNEL_QUALITY_INTERVAL_SEC,
parseTunnelQualityIntervalSeconds,
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
validateTunnelQualityInterval,
} from "@/config/tunnel-quality";
import {
type UpdateReleaseChannel,
getUpdateReleaseChannel,
@@ -157,6 +165,16 @@ const CONFIG_ITEMS: ConfigItem[] = [
"关闭后,前端停止自动刷新,后端停止实时隧道质量探测(全局配置)",
type: "switch",
},
{
key: TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
label: "隧道质量探测间隔",
placeholder: String(DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC),
description:
"设置实时隧道质量检测的执行频率,单位为秒;允许 1–3600 秒,默认 1 秒。",
type: "input",
dependsOn: "monitor_tunnel_quality_enabled",
dependsValue: "true",
},
{
key: "monitor_retention_days",
label: "监控数据保留天数",
@@ -239,6 +257,7 @@ const getInitialConfigs = (): Record<string, string> => {
"cloudflare_secret_key",
"forward_compact_mode",
"monitor_tunnel_quality_enabled",
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
"monitor_retention_days",
"ip",
"panel_domain",
@@ -622,6 +641,19 @@ export default function ConfigPage() {
// 保存配置
const handleSave = async () => {
const intervalValue = configs[TUNNEL_QUALITY_INTERVAL_CONFIG_KEY];
const intervalChanged =
intervalValue !== originalConfigs[TUNNEL_QUALITY_INTERVAL_CONFIG_KEY];
const intervalError = intervalChanged
? validateTunnelQualityInterval(intervalValue || "")
: null;
if (intervalError) {
toast.error(intervalError);
return;
}
setSaving(true);
try {
const changedKeys = Object.keys(configs).filter(
@@ -667,12 +699,22 @@ export default function ConfigPage() {
}),
);
// 如果隧道质量检测开关变更,通知 tunnel-monitor-view
if (changedKeys.includes("monitor_tunnel_quality_enabled")) {
// 如果隧道质量检测配置变更,通知 tunnel-monitor-view
if (
changedKeys.some((key) =>
[
"monitor_tunnel_quality_enabled",
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
].includes(key),
)
) {
window.dispatchEvent(
new CustomEvent("monitorTunnelQualityEnabledChanged", {
detail: {
enabled: configs["monitor_tunnel_quality_enabled"] === "true",
intervalSec: parseTunnelQualityIntervalSeconds(
configs[TUNNEL_QUALITY_INTERVAL_CONFIG_KEY],
),
},
}),
);
@@ -1075,11 +1117,21 @@ export default function ConfigPage() {
case "bg_image":
return renderBgImageUploader();
case "input":
case "input": {
if (isBrandPreviewKey(item.key)) {
return renderBrandAssetUploader(item.key, isChanged);
}
const isTunnelQualityInterval =
item.key === TUNNEL_QUALITY_INTERVAL_CONFIG_KEY;
const intervalValue = isTunnelQualityInterval
? (configs[item.key] ?? String(DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC))
: (configs[item.key] ?? "");
const intervalError =
isTunnelQualityInterval && configs[item.key] !== undefined
? validateTunnelQualityInterval(intervalValue)
: null;
return (
<Input
classNames={{
@@ -1091,14 +1143,30 @@ export default function ConfigPage() {
description={
isCommercialDisabled ? "需商业版授权才能修改此项" : undefined
}
endContent={isTunnelQualityInterval ? "秒" : undefined}
errorMessage={intervalError || undefined}
isDisabled={isCommercialDisabled}
isInvalid={Boolean(intervalError)}
max={
isTunnelQualityInterval
? MAX_TUNNEL_QUALITY_INTERVAL_SEC
: undefined
}
min={
isTunnelQualityInterval
? MIN_TUNNEL_QUALITY_INTERVAL_SEC
: undefined
}
placeholder={item.placeholder}
size="md"
value={configs[item.key] || ""}
step={isTunnelQualityInterval ? 1 : undefined}
type={isTunnelQualityInterval ? "number" : "text"}
value={intervalValue}
variant="bordered"
onChange={(e) => handleConfigChange(item.key, e.target.value)}
/>
);
}
case "switch":
return (
+250 -23
View File
@@ -69,6 +69,7 @@ import {
getNodeList,
pauseForwardService,
resumeForwardService,
resetForwardFlow,
diagnoseForward,
updateForwardOrder,
getConfigByName,
@@ -106,6 +107,7 @@ import { JwtUtil } from "@/utils/jwt";
interface Forward {
id: number;
name: string;
forwardMode?: "agent" | "nftables";
tunnelId: number;
tunnelName: string;
tunnelTrafficRatio?: number;
@@ -129,11 +131,14 @@ interface Forward {
ipSpeedId?: number | null;
ipSpeedLimitName?: string;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
}
interface Tunnel {
id: number;
name: string;
forwardMode?: "agent" | "nftables";
type?: number;
inIp?: string;
inNodeId?: Array<{ nodeId: number }>;
@@ -167,6 +172,8 @@ interface ForwardForm {
ipSpeedId: number | null;
maxConn?: number;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
}
interface ForwardUserGroup {
@@ -224,7 +231,7 @@ const FORWARD_GROUPED_TABLE_COLUMN_CLASS = {
strategy: "w-[100px]",
totalFlow: "w-[120px]",
status: "w-[100px]",
actions: "w-[144px] text-right",
actions: "w-[176px] text-right",
} as const;
const normalizeForwardUserName = (userName?: string): string => {
@@ -598,6 +605,16 @@ const mapForwardApiItems = (items: ForwardApiItem[]): Forward[] => {
typeof forward.proxyProtocol === "number"
? forward.proxyProtocol
: undefined,
proxyProtocolReceive:
typeof forward.proxyProtocolReceive === "number"
? forward.proxyProtocolReceive
: 0,
proxyProtocolSend:
typeof forward.proxyProtocolSend === "number"
? forward.proxyProtocolSend
: typeof forward.proxyProtocol === "number"
? forward.proxyProtocol
: 0,
serviceRunning: forward.status === 1,
}));
};
@@ -748,6 +765,7 @@ const SortableTableRow = ({
handleEdit,
handleDelete,
handleDiagnose,
handleResetFlow,
showAddressModal,
formatFlow,
}: any) => {
@@ -903,6 +921,29 @@ const SortableTableRow = ({
/>
</svg>
</Button>
<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>
<Button
isIconOnly
className="bg-danger/10 text-danger hover:bg-danger/20"
@@ -942,6 +983,7 @@ const SortableCompactTableRow = ({
handleEdit,
handleDelete,
handleDiagnose,
handleResetFlow,
showAddressModal,
hasMultipleAddresses,
formatFlow,
@@ -1128,6 +1170,29 @@ const SortableCompactTableRow = ({
/>
</svg>
</Button>
<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>
<Button
isIconOnly
className="bg-danger/10 text-danger hover:bg-danger/20"
@@ -1273,13 +1338,18 @@ export default function ForwardPage() {
const [modalOpen, setModalOpen] = useState(false);
// isFilterModalOpen removed
const [deleteModalOpen, setDeleteModalOpen] = useState(false);
const [resetFlowModalOpen, setResetFlowModalOpen] = useState(false);
const [addressModalOpen, setAddressModalOpen] = useState(false);
const [diagnosisModalOpen, setDiagnosisModalOpen] = useState(false);
const [isEdit, setIsEdit] = useState(false);
const [submitLoading, setSubmitLoading] = useState(false);
const [deleteLoading, setDeleteLoading] = useState(false);
const [resetFlowLoading, setResetFlowLoading] = useState(false);
const [diagnosisLoading, setDiagnosisLoading] = useState(false);
const [forwardToDelete, setForwardToDelete] = useState<Forward | null>(null);
const [forwardToResetFlow, setForwardToResetFlow] = useState<Forward | null>(
null,
);
const [currentDiagnosisForward, setCurrentDiagnosisForward] =
useState<Forward | null>(null);
const [diagnosisResult, setDiagnosisResult] =
@@ -1335,6 +1405,8 @@ export default function ForwardPage() {
ipSpeedId: null,
maxConn: 0,
proxyProtocol: 0,
proxyProtocolReceive: 0,
proxyProtocolSend: 0,
});
const [inIpTouched, setInIpTouched] = useState(false);
@@ -2126,6 +2198,8 @@ export default function ForwardPage() {
ipMaxConn: 0,
ipSpeedId: null,
proxyProtocol: 0,
proxyProtocolReceive: 0,
proxyProtocolSend: 0,
});
setErrors({});
setModalOpen(true);
@@ -2150,6 +2224,9 @@ export default function ForwardPage() {
ipSpeedId: normalizeSpeedId(forward.ipSpeedId),
maxConn: forward.maxConn ?? 0,
proxyProtocol: forward.proxyProtocol ?? 0,
proxyProtocolReceive: forward.proxyProtocolReceive ?? 0,
proxyProtocolSend:
forward.proxyProtocolSend ?? forward.proxyProtocol ?? 0,
});
setErrors({});
setModalOpen(true);
@@ -2161,6 +2238,46 @@ export default function ForwardPage() {
setDeleteModalOpen(true);
};
const handleResetFlow = (forward: Forward) => {
if ((forward.inFlow || 0) + (forward.outFlow || 0) <= 0) return;
setForwardToResetFlow(forward);
setResetFlowModalOpen(true);
};
const handleResetFlowModalOpenChange = (isOpen: boolean) => {
if (resetFlowLoading) return;
setResetFlowModalOpen(isOpen);
if (!isOpen) {
setForwardToResetFlow(null);
}
};
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);
}
};
// 确认删除规则
const confirmDelete = async () => {
if (!forwardToDelete) return;
@@ -2283,7 +2400,9 @@ export default function ForwardPage() {
ipMaxConn: form.ipMaxConn,
...(isAdmin ? { ipSpeedId: normalizedIPSpeedId } : {}),
maxConn: form.maxConn,
proxyProtocol: form.proxyProtocol,
proxyProtocol: form.proxyProtocolSend,
proxyProtocolReceive: form.proxyProtocolReceive,
proxyProtocolSend: form.proxyProtocolSend,
};
res = await updateForward(updateData);
@@ -2299,7 +2418,9 @@ export default function ForwardPage() {
ipMaxConn: form.ipMaxConn,
...(isAdmin ? { ipSpeedId: normalizedIPSpeedId } : {}),
maxConn: form.maxConn,
proxyProtocol: form.proxyProtocol,
proxyProtocol: form.proxyProtocolSend,
proxyProtocolReceive: form.proxyProtocolReceive,
proxyProtocolSend: form.proxyProtocolSend,
};
res = await createForward(createData);
@@ -3931,7 +4052,7 @@ export default function ForwardPage() {
</div>
</div>
<div className="flex gap-1.5 mt-3">
<div className="grid grid-cols-2 gap-1.5 mt-3">
<Button
className="flex-1 min-h-8"
color="primary"
@@ -3974,6 +4095,32 @@ export default function ForwardPage() {
>
诊断
</Button>
<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>
<Button
className="flex-1 min-h-8"
color="danger"
@@ -4276,7 +4423,7 @@ export default function ForwardPage() {
<TableColumn className="w-[80px]">策略</TableColumn>
<TableColumn className="w-[100px]">用量</TableColumn>
<TableColumn className="w-[80px]">状态</TableColumn>
<TableColumn align="left" className="w-[120px] pl-4">
<TableColumn align="left" className="w-[160px] pl-4">
操作
</TableColumn>
</TableHeader>
@@ -4295,6 +4442,7 @@ export default function ForwardPage() {
handleDelete={handleDelete}
handleDiagnose={handleDiagnose}
handleEdit={handleEdit}
handleResetFlow={handleResetFlow}
handleServiceToggle={handleServiceToggle}
hasMultipleAddresses={hasMultipleAddresses}
selectMode={selectMode}
@@ -4590,6 +4738,7 @@ export default function ForwardPage() {
handleDelete={handleDelete}
handleDiagnose={handleDiagnose}
handleEdit={handleEdit}
handleResetFlow={handleResetFlow}
handleServiceToggle={
handleServiceToggle
}
@@ -4966,25 +5115,50 @@ export default function ForwardPage() {
setForm((prev) => ({ ...prev, ipMaxConn: value }));
}}
/>
<Select
description="启用 PROXY protocol,用于透传客户端真实 IP"
label="Proxy Protocol"
placeholder="禁用"
selectedKeys={[String(form.proxyProtocol || 0)]}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as string;
<div className="grid grid-cols-1 gap-4 md:grid-cols-2">
<Select
description="入口监听接收 PROXY protocol,用于读取上游传入的真实客户端 IP。"
label="Proxy Protocol 接收"
placeholder="禁用"
selectedKeys={[
String(form.proxyProtocolReceive || 0),
]}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as string;
setForm((prev) => ({
...prev,
proxyProtocol: Number(selectedKey),
}));
}}
>
<SelectItem key="0">禁用</SelectItem>
<SelectItem key="1">Version 1</SelectItem>
<SelectItem key="2">Version 2</SelectItem>
</Select>
setForm((prev) => ({
...prev,
proxyProtocolReceive: Number(selectedKey),
}));
}}
>
<SelectItem key="0">禁用</SelectItem>
<SelectItem key="1">Version 1</SelectItem>
<SelectItem key="2">Version 2</SelectItem>
</Select>
<Select
description="连接目标地址时发送 PROXY protocol,用于向下游透传客户端真实 IP。"
label="Proxy Protocol 发送"
placeholder="禁用"
selectedKeys={[String(form.proxyProtocolSend || 0)]}
variant="bordered"
onSelectionChange={(keys) => {
const selectedKey = Array.from(keys)[0] as string;
const proxyProtocolSend = Number(selectedKey);
setForm((prev) => ({
...prev,
proxyProtocol: proxyProtocolSend,
proxyProtocolSend,
}));
}}
>
<SelectItem key="0">禁用</SelectItem>
<SelectItem key="1">Version 1</SelectItem>
<SelectItem key="2">Version 2</SelectItem>
</Select>
</div>
{isAdmin && (
<Select
label="规则限速"
@@ -5121,6 +5295,59 @@ export default function ForwardPage() {
</ModalContent>
</Modal>
{/* 规则流量清零确认模态框 */}
<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={handleResetFlowModalOpenChange}
>
<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">
&quot;{forwardToResetFlow?.name}&quot;
</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>
{/* 地址列表弹窗 */}
<Modal
classNames={{
+316 -10
View File
@@ -109,6 +109,17 @@ interface Node {
socks?: number; // 0 关 1 开
status: number;
isRemote?: number;
forwardMode?: "agent" | "nftables";
sshConfig?: {
host?: string;
port?: number;
username?: string;
authType?: "password" | "private_key";
password?: string;
privateKey?: string;
passphrase?: string;
sudoMode?: "none" | "sudo" | "sudo_su";
} | null;
remoteUrl?: string;
syncError?: string;
connectionStatus: "online" | "offline";
@@ -132,6 +143,15 @@ interface NodeForm {
udpListenAddr: string;
interfaceName: string;
extraIPs: string;
forwardMode: "agent" | "nftables";
sshHost: string;
sshPort: string;
sshUsername: string;
sshAuthType: "password" | "private_key";
sshPassword: string;
sshPrivateKey: string;
sshPassphrase: string;
sshSudoMode: "none" | "sudo" | "sudo_su";
http: number; // 0 关 1 开
tls: number; // 0 关 1 开
socks: number; // 0 关 1 开
@@ -344,11 +364,25 @@ export default function NodePage() {
udpListenAddr: "[::]",
interfaceName: "",
extraIPs: "",
forwardMode: "agent",
sshHost: "",
sshPort: "22",
sshUsername: "",
sshAuthType: "private_key",
sshPassword: "",
sshPrivateKey: "",
sshPassphrase: "",
sshSudoMode: "none",
http: 0,
tls: 0,
socks: 0,
});
const [errors, setErrors] = useState<Record<string, string>>({});
const isNftablesMode = form.forwardMode === "nftables";
const protocolControlsDisabled = protocolDisabled || isNftablesMode;
const protocolControlsReason = isNftablesMode
? "nftables 模式不支持 agent 协议开关"
: protocolDisabledReason || "等待节点上线后再设置";
const [selectMode, setSelectMode] = useState(false);
const [selectedIds, setSelectedIds] = useState<Set<number>>(new Set());
@@ -856,6 +890,27 @@ export default function NodePage() {
newErrors.port = portValidation.error || "端口格式错误";
}
if (form.forwardMode === "nftables") {
if (!form.sshHost.trim()) {
newErrors.sshHost = "请输入 SSH 主机";
}
const sshPort = Number(form.sshPort);
if (!Number.isInteger(sshPort) || sshPort < 1 || sshPort > 65535) {
newErrors.sshPort = "SSH 端口必须为 1-65535";
}
if (!form.sshUsername.trim()) {
newErrors.sshUsername = "请输入 SSH 用户名";
}
if (form.sshAuthType === "password") {
if (!isEdit && !form.sshPassword.trim()) {
newErrors.sshPassword = "请输入 SSH 密码";
}
} else if (!isEdit && !form.sshPrivateKey.trim()) {
newErrors.sshPrivateKey = "请输入 SSH 私钥";
}
}
setErrors(newErrors);
return Object.keys(newErrors).length === 0;
@@ -898,15 +953,28 @@ export default function NodePage() {
udpListenAddr: node.udpListenAddr || "[::]",
interfaceName: (node as any).interfaceName || "",
extraIPs: node.extraIPs || "",
forwardMode: node.forwardMode === "nftables" ? "nftables" : "agent",
sshHost: node.sshConfig?.host || normalizedHost || legacy,
sshPort: String(node.sshConfig?.port || 22),
sshUsername: node.sshConfig?.username || "",
sshAuthType: node.sshConfig?.authType || "private_key",
sshPassword: node.sshConfig?.password || "",
sshPrivateKey: node.sshConfig?.privateKey || "",
sshPassphrase: node.sshConfig?.passphrase || "",
sshSudoMode: node.sshConfig?.sudoMode || "none",
http: typeof node.http === "number" ? node.http : 1,
tls: typeof node.tls === "number" ? node.tls : 1,
socks: typeof node.socks === "number" ? node.socks : 1,
});
const offline = node.connectionStatus !== "online";
setProtocolDisabled(offline);
setProtocolDisabled(offline || node.forwardMode === "nftables");
setProtocolDisabledReason(
offline ? "节点未在线,等待节点上线后再设置" : "",
node.forwardMode === "nftables"
? "nftables 模式不支持 agent 协议开关"
: offline
? "节点未在线,等待节点上线后再设置"
: "",
);
setDialogVisible(true);
};
@@ -1166,12 +1234,33 @@ export default function NodePage() {
try {
const apiCall = isEdit ? updateNode : createNode;
const { serverHost, ...rest } = form;
const sshConfig =
form.forwardMode === "nftables"
? {
host: form.sshHost.trim(),
port: Number(form.sshPort || 22),
username: form.sshUsername.trim(),
authType: form.sshAuthType,
password:
form.sshAuthType === "password"
? form.sshPassword.trim()
: undefined,
privateKey:
form.sshAuthType === "private_key"
? form.sshPrivateKey
: undefined,
passphrase: form.sshPassphrase,
sudoMode: form.sshSudoMode,
}
: null;
const data = {
...rest,
remark: form.remark.trim(),
expiryTime: form.expiryTime,
renewalCycle: form.renewalCycle,
extraIPs: form.extraIPs,
forwardMode: form.forwardMode,
sshConfig,
serverIp:
form.serverIpV4?.trim() ||
form.serverIpV6?.trim() ||
@@ -1206,6 +1295,8 @@ export default function NodePage() {
tcpListenAddr: form.tcpListenAddr,
udpListenAddr: form.udpListenAddr,
interfaceName: form.interfaceName,
forwardMode: form.forwardMode,
sshConfig,
http: form.http,
tls: form.tls,
socks: form.socks,
@@ -1242,6 +1333,15 @@ export default function NodePage() {
udpListenAddr: "[::]",
interfaceName: "",
extraIPs: "",
forwardMode: "agent",
sshHost: "",
sshPort: "22",
sshUsername: "",
sshAuthType: "password",
sshPassword: "",
sshPrivateKey: "",
sshPassphrase: "",
sshSudoMode: "none",
http: 0,
tls: 0,
socks: 0,
@@ -2581,6 +2681,43 @@ export default function NodePage() {
}
/>
<Select
label="转发模式"
selectedKeys={[form.forwardMode]}
variant="bordered"
onSelectionChange={(keys) => {
const selected = Array.from(keys)[0] as
| NodeForm["forwardMode"]
| undefined;
setForm((prev) => ({
...prev,
forwardMode: selected || "agent",
}));
}}
>
<SelectItem key="agent" textValue="agent">
agent 模式
</SelectItem>
<SelectItem key="nftables" textValue="nftables">
nftables 模式
</SelectItem>
</Select>
{form.forwardMode === "nftables" ? (
<Alert
color="warning"
description="nftables 模式仅支持纯转发,不支持隧道、流量控制和 agent 安装,规则将由面板通过 SSH 下发。"
variant="flat"
/>
) : (
<Alert
color="primary"
description="agent 模式会继续使用节点 agent 执行转发与管理操作。"
variant="flat"
/>
)}
{/* 高级配置 */}
<Accordion className="px-0" variant="light">
<AccordionItem
@@ -2594,6 +2731,177 @@ export default function NodePage() {
}
>
<div className="space-y-4 pb-2">
{form.forwardMode === "nftables" ? (
<div className="space-y-4 rounded-xl border border-warning-200 bg-warning-50/60 p-4 dark:border-warning-500/30 dark:bg-warning-950/20">
<div>
<div className="text-sm font-medium text-warning-700 dark:text-warning-300">
SSH 配置
</div>
<div className="text-xs text-default-500">
面板会通过 SSH 下发 nftables 规则,不会安装 agent。
</div>
</div>
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
<Input
errorMessage={errors.sshHost}
isInvalid={!!errors.sshHost}
label="SSH 主机"
placeholder="例如: 192.0.2.10"
value={form.sshHost}
variant="bordered"
onChange={(e) =>
setForm((prev) => ({
...prev,
sshHost: e.target.value,
}))
}
/>
<Input
errorMessage={errors.sshPort}
isInvalid={!!errors.sshPort}
label="SSH 端口"
placeholder="22"
value={form.sshPort}
variant="bordered"
onChange={(e) =>
setForm((prev) => ({
...prev,
sshPort: e.target.value,
}))
}
/>
</div>
<div className="grid grid-cols-1 md:grid-cols-2 gap-4">
<Input
errorMessage={errors.sshUsername}
isInvalid={!!errors.sshUsername}
label="SSH 用户名"
placeholder="例如: root"
value={form.sshUsername}
variant="bordered"
onChange={(e) =>
setForm((prev) => ({
...prev,
sshUsername: e.target.value,
}))
}
/>
<Select
label="SSH 认证方式"
selectedKeys={[form.sshAuthType]}
variant="bordered"
onSelectionChange={(keys) => {
const selected = Array.from(keys)[0] as
| "password"
| "private_key"
| undefined;
setForm((prev) => ({
...prev,
sshAuthType: selected || "private_key",
}));
}}
>
<SelectItem
key="private_key"
textValue="private_key"
>
私钥
</SelectItem>
<SelectItem key="password" textValue="password">
密码
</SelectItem>
</Select>
</div>
{form.sshAuthType === "password" ? (
<Input
errorMessage={errors.sshPassword}
isInvalid={!!errors.sshPassword}
label="SSH 密码"
placeholder="请输入 SSH 密码"
type="password"
value={form.sshPassword}
variant="bordered"
onChange={(e) =>
setForm((prev) => ({
...prev,
sshPassword: e.target.value,
}))
}
/>
) : (
<Textarea
errorMessage={errors.sshPrivateKey}
isInvalid={!!errors.sshPrivateKey}
label="SSH 私钥"
minRows={6}
placeholder="粘贴 PEM 格式私钥内容"
value={form.sshPrivateKey}
variant="bordered"
onChange={(e) =>
setForm((prev) => ({
...prev,
sshPrivateKey: e.target.value,
}))
}
/>
)}
<Input
label="SSH 私钥密码短语"
placeholder="可选"
type="password"
value={form.sshPassphrase}
variant="bordered"
onChange={(e) =>
setForm((prev) => ({
...prev,
sshPassphrase: e.target.value,
}))
}
/>
<Select
label="sudo 提权方式"
selectedKeys={[form.sshSudoMode]}
variant="bordered"
onSelectionChange={(keys) => {
const selected = Array.from(keys)[0] as
| "none"
| "sudo"
| "sudo_su"
| undefined;
setForm((prev) => ({
...prev,
sshSudoMode: selected || "none",
}));
}}
>
<SelectItem key="none" textValue="none">
无需提权
</SelectItem>
<SelectItem key="sudo" textValue="sudo">
sudo
</SelectItem>
<SelectItem key="sudo_su" textValue="sudo_su">
sudo su
</SelectItem>
</Select>
</div>
) : (
<Alert
color="primary"
description="agent 模式下无需填写 SSH 配置;节点安装 agent 后会自行接管转发。"
variant="flat"
/>
)}
<Input
description="用于多IP服务器指定使用那个IP请求远程地址,不懂的默认为空就行"
errorMessage={errors.interfaceName}
@@ -2677,18 +2985,16 @@ export default function NodePage() {
<div className="text-xs text-default-500 mb-2">
开启开关以屏蔽对应协议
</div>
{protocolDisabled && (
{protocolControlsDisabled && (
<Alert
className="mb-2"
color="warning"
description={
protocolDisabledReason || "等待节点上线后再设置"
}
description={protocolControlsReason}
variant="flat"
/>
)}
<div
className={`grid grid-cols-1 sm:grid-cols-3 gap-3 bg-content1/30 dark:bg-content1/20 p-3 rounded-md border border-divider ${protocolDisabled ? "opacity-70" : ""}`}
className={`grid grid-cols-1 sm:grid-cols-3 gap-3 bg-content1/30 dark:bg-content1/20 p-3 rounded-md border border-divider ${protocolControlsDisabled ? "opacity-70" : ""}`}
>
{/* HTTP tile */}
<div className="px-3 py-3 rounded-lg bg-content1/55 dark:bg-content1/35 border border-divider hover:border-primary-200 dark:hover:border-primary-500/30 transition-colors">
@@ -2715,7 +3021,7 @@ export default function NodePage() {
禁用/启用
</div>
<Switch
isDisabled={protocolDisabled}
isDisabled={protocolControlsDisabled}
isSelected={form.http === 1}
size="sm"
onValueChange={(v) =>
@@ -2762,7 +3068,7 @@ export default function NodePage() {
禁用/启用
</div>
<Switch
isDisabled={protocolDisabled}
isDisabled={protocolControlsDisabled}
isSelected={form.tls === 1}
size="sm"
onValueChange={(v) =>
@@ -2801,7 +3107,7 @@ export default function NodePage() {
禁用/启用
</div>
<Switch
isDisabled={protocolDisabled}
isDisabled={protocolControlsDisabled}
isSelected={form.socks === 1}
size="sm"
onValueChange={(v) =>
@@ -2,6 +2,8 @@ import type {
MonitorTunnelApiItem,
TunnelMetricApiItem,
TunnelQualityApiItem,
TunnelQualityCandidateHopApiItem,
TunnelQualityChainDetailsApiItem,
TunnelQualityHopApiItem,
} from "@/api/types";
@@ -53,12 +55,17 @@ import {
TableRow,
TableCell,
} from "@/shadcn-bridge/heroui/table";
import {
DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC,
parseTunnelQualityIntervalSeconds,
TUNNEL_QUALITY_INTERVAL_CONFIG_KEY,
tunnelQualityIntervalLabel,
} from "@/config/tunnel-quality";
interface TunnelMonitorViewProps {
viewMode?: "list" | "grid";
}
const QUALITY_POLL_INTERVAL = 1_000; // 1 second
const MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY =
"monitor_tunnel_quality_enabled";
const MONITOR_TUNNEL_QUALITY_ENABLED_EVENT =
@@ -463,18 +470,59 @@ const TrafficChartCard = React.memo(function TrafficChartCard({
);
});
function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
if (!hopsStr) return null;
let hops: TunnelQualityHopApiItem[] = [];
const parseTunnelQualityChainDetails = (
raw?: string,
): TunnelQualityChainDetailsApiItem => {
if (!raw) return {};
try {
hops = JSON.parse(hopsStr);
const parsed: unknown = JSON.parse(raw);
// Backward-compatible with historical rows that stored the primary path
// directly as a JSON array.
if (Array.isArray(parsed)) {
return { primaryPath: parsed as TunnelQualityHopApiItem[] };
}
if (parsed && typeof parsed === "object") {
return parsed as TunnelQualityChainDetailsApiItem;
}
} catch {
return null;
return {};
}
if (!Array.isArray(hops) || hops.length === 0) return null;
return {};
};
const candidateRoleLabel = (role: string) => {
switch (role) {
case "entry":
return "入口";
case "middle":
return "中转";
case "exit":
return "出口";
case "target":
return "测试目标";
default:
return role;
}
};
const ForwardingChainTopology = React.memo(function ForwardingChainTopology({
hopsStr,
}: {
hopsStr?: string;
}) {
const details = useMemo(
() => parseTunnelQualityChainDetails(hopsStr),
[hopsStr],
);
const primaryHops = details.primaryPath ?? [];
const alternativeHops = useMemo(
() => (details.candidateHops ?? []).filter((hop) => !hop.selected),
[details.candidateHops],
);
if (primaryHops.length === 0 && alternativeHops.length === 0) return null;
return (
<Card className="border border-divider/60 shadow-sm transition-shadow bg-gradient-to-br from-background to-default-50/50 mt-4">
@@ -485,61 +533,131 @@ function ForwardingChainTopology({ hopsStr }: { hopsStr?: string }) {
</h3>
</CardHeader>
<CardBody className="py-2 px-4 pb-4">
<div className="flex items-center overflow-x-auto pb-2 py-2">
{hops.map((hop, index) => {
const hasError = hop.latency < 0 || hop.loss > 0;
const colorClass =
hop.latency < 0
? "text-danger"
: hop.loss > 0
? "text-warning"
: "text-success";
const borderColor = hasError ? "border-danger" : "";
{primaryHops.length > 0 ? (
<div className="flex items-center overflow-x-auto pb-2 py-2">
{primaryHops.map((hop, index) => {
const hasError = hop.latency < 0 || hop.loss > 0;
const colorClass =
hop.latency < 0
? "text-danger"
: hop.loss > 0
? "text-warning"
: "text-success";
const borderColor = hasError ? "border-danger" : "";
return (
<React.Fragment key={index}>
{index === 0 && (
return (
<React.Fragment
key={`${hop.fromNodeId}-${hop.toNodeId}-${index}`}
>
{index === 0 ? (
<Chip
className="shrink-0 font-mono shadow-sm"
size="sm"
variant="flat"
>
{hop.fromNodeName}
</Chip>
) : null}
<div className="flex flex-col items-center justify-center min-w-[70px] mx-1 shrink-0 relative">
<span
className={`text-[10px] font-mono leading-none mb-1 ${colorClass}`}
>
{hop.latency >= 0
? `${hop.latency.toFixed(0)}ms`
: "超时"}
</span>
<div
className={`h-[2px] w-full relative flex items-center justify-end bg-default-200 ${hop.latency < 0 ? "!bg-danger" : ""}`}
>
<ArrowRight
className={`w-3.5 h-3.5 absolute -right-2 ${colorClass} bg-background rounded-full p-[1px] z-10`}
/>
</div>
<span
className={`text-[10px] font-mono leading-none mt-1.5 ${hop.loss > 0 ? "text-warning" : "text-default-400"}`}
>
{hop.loss.toFixed(0)}% 丢包
</span>
</div>
<Chip
className="shrink-0 font-mono shadow-sm"
className={`shrink-0 font-mono shadow-sm ${borderColor}`}
size="sm"
variant="flat"
>
{hop.fromNodeName}
{hop.toNodeName}
</Chip>
)}
<div className="flex flex-col items-center justify-center min-w-[70px] mx-1 shrink-0 relative">
<span
className={`text-[10px] font-mono leading-none mb-1 ${colorClass}`}
>
{hop.latency >= 0 ? `${hop.latency.toFixed(0)}ms` : "超时"}
</span>
<div
className={`h-[2px] w-full relative flex items-center justify-end bg-default-200 ${hop.latency < 0 ? "!bg-danger" : ""}`}
>
<ArrowRight
className={`w-3.5 h-3.5 absolute -right-2 ${colorClass} bg-background rounded-full p-[1px] z-10`}
/>
</div>
<span
className={`text-[10px] font-mono leading-none mt-1.5 ${hop.loss > 0 ? "text-warning" : "text-default-400"}`}
>
{hop.loss.toFixed(0)}% 丢包
</span>
</div>
<Chip
className={`shrink-0 font-mono shadow-sm ${borderColor}`}
size="sm"
variant="flat"
>
{hop.toNodeName}
</Chip>
</React.Fragment>
);
})}
</div>
</React.Fragment>
);
})}
</div>
) : null}
{alternativeHops.length > 0 ? (
<div
className={`${primaryHops.length > 0 ? "mt-3 border-t border-divider/60 pt-3" : ""}`}
>
<div className="mb-2 flex items-center justify-between gap-2">
<span className="text-xs font-semibold text-default-600">
备选节点实时延迟
</span>
<span className="text-[10px] text-default-400">
{alternativeHops.length} 条候选链路
</span>
</div>
<div className="grid max-h-72 grid-cols-1 gap-2 overflow-y-auto pr-1 md:grid-cols-2">
{alternativeHops.map((hop) => (
<CandidateHopLatency
key={`${hop.hopIndex}-${hop.fromNodeId}-${hop.toRole}-${hop.toNodeId || hop.toNodeName}`}
hop={hop}
/>
))}
</div>
</div>
) : null}
</CardBody>
</Card>
);
});
function CandidateHopLatency({
hop,
}: {
hop: TunnelQualityCandidateHopApiItem;
}) {
const hasError = Boolean(hop.errorMessage) || hop.latency < 0;
return (
<div className="rounded-xl border border-divider/60 bg-default-50/50 px-3 py-2.5">
<div className="flex items-center gap-2 text-xs">
<span className="min-w-0 truncate font-medium text-foreground">
{hop.fromNodeName}
</span>
<ArrowRight className="h-3.5 w-3.5 shrink-0 text-default-400" />
<span className="min-w-0 truncate font-medium text-foreground">
{hop.toNodeName}
</span>
<span className="ml-auto shrink-0">
<LatencyDisplay value={hop.latency} />
</span>
</div>
<div className="mt-1.5 flex items-center gap-2 text-[10px] text-default-400">
<span>
{candidateRoleLabel(hop.fromRole)} → {candidateRoleLabel(hop.toRole)}
</span>
{hasError ? (
<span className="ml-auto truncate text-danger">
{hop.errorMessage || "探测失败"}
</span>
) : (
<span
className={`ml-auto font-mono ${hop.loss > 0 ? "text-warning" : "text-default-500"}`}
>
{hop.loss.toFixed(0)}% 丢包
</span>
)}
</div>
</div>
);
}
export function TunnelMonitorView({
@@ -562,6 +680,9 @@ export function TunnelMonitorView({
const qualityTimerRef = useRef<number | null>(null);
const [monitorTunnelQualityEnabled, setMonitorTunnelQualityEnabled] =
useState(true);
const [tunnelQualityIntervalSec, setTunnelQualityIntervalSec] = useState(
DEFAULT_TUNNEL_QUALITY_INTERVAL_SEC,
);
// Detail view state
const [detailTunnelId, setDetailTunnelId] = useState<number | null>(null);
@@ -618,26 +739,28 @@ export function TunnelMonitorView({
}
}, []);
const loadMonitorTunnelQualityEnabled = useCallback(async () => {
try {
const response = await getConfigByName(
MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY,
);
const loadTunnelQualityConfig = useCallback(async () => {
const [enabledResponse, intervalResponse] = await Promise.all([
getConfigByName(MONITOR_TUNNEL_QUALITY_ENABLED_CONFIG_KEY).catch(
() => null,
),
getConfigByName(TUNNEL_QUALITY_INTERVAL_CONFIG_KEY).catch(() => null),
]);
setMonitorTunnelQualityEnabled(
typeof response.data?.value === "string"
? response.data.value === "true"
: true,
);
} catch {
setMonitorTunnelQualityEnabled(true);
}
setMonitorTunnelQualityEnabled(
typeof enabledResponse?.data?.value === "string"
? enabledResponse.data.value === "true"
: true,
);
setTunnelQualityIntervalSec(
parseTunnelQualityIntervalSeconds(intervalResponse?.data?.value),
);
}, []);
useEffect(() => {
void loadTunnels();
void loadMonitorTunnelQualityEnabled();
}, [loadMonitorTunnelQualityEnabled, loadTunnels]);
void loadTunnelQualityConfig();
}, [loadTunnelQualityConfig, loadTunnels]);
useEffect(() => {
const timer = window.setInterval(() => {
@@ -649,8 +772,11 @@ export function TunnelMonitorView({
useEffect(() => {
const handleMonitorTunnelQualityEnabledChanged = (event: Event) => {
const enabled = (event as CustomEvent<{ enabled?: boolean }>).detail
?.enabled;
const detail = (
event as CustomEvent<{ enabled?: boolean; intervalSec?: number }>
).detail;
const enabled = detail?.enabled;
const intervalSec = detail?.intervalSec;
if (typeof enabled === "boolean") {
setMonitorTunnelQualityEnabled(enabled);
@@ -658,7 +784,12 @@ export function TunnelMonitorView({
setQualityLoading(false);
}
} else {
void loadMonitorTunnelQualityEnabled();
void loadTunnelQualityConfig();
}
if (typeof intervalSec === "number") {
setTunnelQualityIntervalSec(
parseTunnelQualityIntervalSeconds(String(intervalSec)),
);
}
};
@@ -673,7 +804,7 @@ export function TunnelMonitorView({
handleMonitorTunnelQualityEnabledChanged as EventListener,
);
};
}, [loadMonitorTunnelQualityEnabled]);
}, [loadTunnelQualityConfig]);
useEffect(() => {
if (tunnels.length > 0 && !initialHistoryFetched.current) {
@@ -726,7 +857,7 @@ export function TunnelMonitorView({
}
}, [tunnels]);
// --- Load quality snapshots (auto-polling every 10s) ---
// --- Load quality snapshots using the configured probe interval ---
const loadQuality = useCallback(async (options?: { silent?: boolean }) => {
const silent = options?.silent ?? false;
@@ -788,7 +919,7 @@ export function TunnelMonitorView({
qualityTimerRef.current = window.setInterval(() => {
void loadQuality({ silent: true });
}, QUALITY_POLL_INTERVAL);
}, tunnelQualityIntervalSec * 1000);
return () => {
if (qualityTimerRef.current) {
@@ -796,7 +927,7 @@ export function TunnelMonitorView({
qualityTimerRef.current = null;
}
};
}, [loadQuality, monitorTunnelQualityEnabled]);
}, [loadQuality, monitorTunnelQualityEnabled, tunnelQualityIntervalSec]);
// --- Load quality history for detail chart ---
const loadQualityHistory = useCallback(
@@ -1065,7 +1196,7 @@ export function TunnelMonitorView({
{monitorTunnelQualityEnabled ? (
<>
<LiveDot />
<span>自动探测中(每秒测试,30秒上报)</span>
<span>{`自动探测中(${tunnelQualityIntervalLabel(tunnelQualityIntervalSec)}测试)`}</span>
</>
) : (
<>
@@ -1129,7 +1260,7 @@ export function TunnelMonitorView({
{monitorTunnelQualityEnabled ? (
<>
<LiveDot />
<span>每秒探测 · 更新于 {lastQualityUpdate}</span>
<span>{`${tunnelQualityIntervalLabel(tunnelQualityIntervalSec)}探测 · 更新于 ${lastQualityUpdate}`}</span>
</>
) : (
<>
+32 -3
View File
@@ -122,6 +122,7 @@ interface Tunnel {
inx?: number;
name: string;
type: number; // 1: 端口转发, 2: 隧道转发
forwardMode?: "agent" | "nftables";
inNodeId: ChainTunnel[]; // 入口节点列表
outNodeId?: ChainTunnel[]; // 出口节点列表
chainNodes?: ChainTunnel[][]; // 转发链节点列表,二维数组
@@ -150,6 +151,7 @@ interface Node {
id: number;
name: string;
status: number; // 1: 在线, 0: 离线
forwardMode?: "agent" | "nftables";
serverIp?: string;
serverIpV4?: string;
serverIpV6?: string;
@@ -160,6 +162,7 @@ interface TunnelForm {
id?: number;
name: string;
type: number;
forwardMode?: "agent" | "nftables";
inNodeId: ChainTunnel[];
outNodeId?: ChainTunnel[];
chainNodes?: ChainTunnel[][]; // 转发链节点列表,二维数组,外层是跳数,内层是该跳的节点
@@ -172,6 +175,23 @@ interface TunnelForm {
status: number;
}
type TunnelForwardMode = NonNullable<TunnelForm["forwardMode"]>;
const getTunnelForwardMode = (
inNodeId: ChainTunnel[],
nodes: Node[],
): TunnelForwardMode =>
inNodeId.some((item) => {
const node = nodes.find((candidate) => candidate.id === item.nodeId);
return node?.forwardMode === "nftables";
})
? "nftables"
: "agent";
const createTypedTunnelFormDefaults = (): TunnelForm =>
createTunnelFormDefaults() as TunnelForm;
interface BatchProgressState {
active: boolean;
label: string;
@@ -416,7 +436,7 @@ export default function TunnelPage() {
};
// 表单状态
const [form, setForm] = useState<TunnelForm>(createTunnelFormDefaults());
const [form, setForm] = useState<TunnelForm>(createTypedTunnelFormDefaults());
// 表单验证错误
const [errors, setErrors] = useState<{ [key: string]: string }>({});
@@ -564,7 +584,14 @@ export default function TunnelPage() {
// 表单验证
const validateForm = (): boolean => {
const newErrors = validateTunnelForm(form, nodes, isEdit);
const newErrors = validateTunnelForm(
{
...form,
forwardMode: getTunnelForwardMode(form.inNodeId, nodes),
},
nodes,
isEdit,
);
setErrors(newErrors);
@@ -574,7 +601,7 @@ export default function TunnelPage() {
// 新增隧道
const handleAdd = () => {
setIsEdit(false);
setForm(createTunnelFormDefaults());
setForm(createTypedTunnelFormDefaults());
setErrors({});
setModalOpen(true);
};
@@ -588,6 +615,7 @@ export default function TunnelPage() {
id: tunnel.id,
name: tunnel.name,
type: tunnel.type,
forwardMode: tunnel.forwardMode === "nftables" ? "nftables" : "agent",
inNodeId: tunnel.inNodeId || [],
outNodeId: tunnel.outNodeId || [],
chainNodes: tunnel.chainNodes || [],
@@ -885,6 +913,7 @@ export default function TunnelPage() {
const data = {
...form,
forwardMode: getTunnelForwardMode(form.inNodeId, nodes),
inIp: inIpString,
outNodeId: cleanedOutNodeId,
chainNodes: cleanedChainNodes,
+24
View File
@@ -8,6 +8,7 @@ interface TunnelFormInput {
inNodeId: TunnelChainNode[];
outNodeId?: TunnelChainNode[];
trafficRatio: number;
forwardMode?: "agent" | "nftables";
probeTargetHost?: string;
probeTargetPort?: number;
}
@@ -15,6 +16,7 @@ interface TunnelFormInput {
interface TunnelNodeInput {
id: number;
status: number;
forwardMode?: "agent" | "nftables";
}
const isValidProbeIPv4 = (host: string) => {
@@ -116,6 +118,7 @@ export const createTunnelFormDefaults = () => {
chainNodes: [],
flow: 1,
trafficRatio: 1.0,
forwardMode: "agent",
inIp: "",
ipPreference: "",
probeTargetHost: "",
@@ -183,6 +186,10 @@ export const validateTunnelForm = (
}
if (form.type === 2) {
if (form.forwardMode === "nftables") {
errors.type = "nftables 节点仅支持端口转发";
}
if (!form.outNodeId || form.outNodeId.length === 0) {
errors.outNodeId = "请至少选择一个出口节点";
} else {
@@ -208,6 +215,23 @@ export const validateTunnelForm = (
}
}
if (form.forwardMode === "nftables") {
const nftEntryNodes = (form.inNodeId || []).filter((item) => {
const node = nodes.find((n) => n.id === item.nodeId);
return node?.forwardMode === "nftables";
});
if (nftEntryNodes.length > 0) {
if (form.inNodeId.length !== 1) {
errors.inNodeId = "nftables 节点仅支持单入口隧道";
}
if ((form.outNodeId || []).length > 0) {
errors.outNodeId = "nftables 节点不支持出口节点配置";
}
}
}
return errors;
};