Compare commits

...

14 Commits

Author SHA1 Message Date
sagitchu 6d13ebd6e1 fix(agent): harden Alpine OpenRC installation 2026-08-08 11:43:33 +08:00
sagit 538e64ffc0 Fix modal scroll position jumps (#542)
Preserve page scroll positions while Radix modals acquire focus and forward the dialog overlay ref correctly.
2026-08-07 22:38:02 +08:00
sagit 0b23d6f7d7 fix(agent): retire replaced tunnel sessions (#541) 2026-08-07 14:35:30 +08:00
sagit 9e6f80019d feat(monitor): show backup nodes in topology (#538) 2026-08-03 14:14:01 +08:00
sagit a8fd01d4d8 fix(node): allow IPv6-only addresses (#537) 2026-08-03 11:05:04 +08:00
sagit cbe2fc492e feat(monitor): show backup tunnel latencies (#535)
Closes #508
2026-08-03 10:15:22 +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
63 changed files with 4625 additions and 475 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、`curl` 和 CA 证书,并使用 OpenRC 注册、启动和管理 `flux_agent` 服务;其他受支持的 Linux 发行版继续使用 systemd。
**安装过程中会提示输入:**
- **服务器地址**: 面板端的通信地址(通常是 `http://<面板IP>:<后端端口>`,例如 `http://1.2.3.4:6365`)。
- **密钥**: 刚才在面板中获取的节点密钥。
@@ -77,7 +85,8 @@ curl -L https://raw.githubusercontent.com/Sagit-chu/flux-panel/main/install.sh -
### 3. 验证安装
安装完成后,服务会自动启动。
- 查看状态: `systemctl status flux_agent`
- systemd 查看状态: `systemctl status flux_agent`
- Alpine/OpenRC 查看状态: `rc-service flux_agent status`
- 回到面板 **节点管理** 页面,该节点状态应显示为 **在线**。
---
@@ -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,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
}
@@ -772,6 +774,9 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
if ctx == nil {
ctx = context.Background()
}
if payload, handled, err := h.diagnoseNftablesForwardRuntime(forward); handled || err != nil {
return payload, err
}
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
if err != nil {
return nil, err
@@ -787,6 +792,111 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
return payload, nil
}
func (h *Handler) diagnoseNftablesForwardRuntime(forward *forwardRecord) (map[string]interface{}, bool, error) {
if forward == nil {
return nil, false, errForwardNotFound
}
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
if err != nil {
return nil, false, err
}
if !nftMode {
return nil, false, nil
}
if len(entryNodeIDs) == 0 {
return nil, true, errors.New("nftables 转发缺少入口节点")
}
targets, err := resolveDiagnosisTargets(forward.RemoteAddr)
if err != nil {
return nil, true, err
}
results, err := h.buildNftablesForwardDiagnosisResults(forward, entryNodeIDs[0], targets)
if err != nil {
return nil, true, err
}
payload := map[string]interface{}{
"forwardName": forward.Name,
"timestamp": time.Now().UnixMilli(),
"results": results,
}
return payload, true, nil
}
func (h *Handler) buildNftablesForwardDiagnosisResults(forward *forwardRecord, nodeID int64, targets []diagnosisTarget) ([]map[string]interface{}, error) {
if h == nil || h.repo == nil {
return nil, errors.New("handler not initialized")
}
node, err := h.getNodeRecord(nodeID)
if err != nil {
return nil, err
}
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
if err != nil {
return nil, err
}
var binding *model.NftRuleBinding
for i := range bindings {
if bindings[i].ForwardID == forward.ID {
binding = &bindings[i]
break
}
}
target := diagnosisTarget{}
if len(targets) > 0 {
target = targets[0]
}
status := "missing"
message := "nftables 规则未下发"
success := false
inPort := 0
protocols := ""
targetAddr := strings.TrimSpace(forward.RemoteAddr)
ruleHash := ""
if binding != nil {
status = strings.ToLower(strings.TrimSpace(binding.Status))
inPort = binding.InPort
protocols = strings.TrimSpace(binding.Protocols)
targetAddr = strings.TrimSpace(binding.TargetAddr)
ruleHash = strings.TrimSpace(binding.RuleHash)
if status == "" {
status = "pending"
}
if status == runtimenft.StatusApplied {
success = true
message = "nftables 规则已下发"
} else if strings.TrimSpace(binding.LastError) != "" {
message = binding.LastError
} else {
message = "nftables 规则未完成下发"
}
}
packetLoss := 100
if success {
packetLoss = 0
}
result := map[string]interface{}{
"success": success,
"nodeName": node.Name,
"nodeId": strconv.FormatInt(nodeID, 10),
"targetIp": target.IP,
"targetPort": target.Port,
"description": fmt.Sprintf("nftables规则(%s)->目标(%s)", node.Name, defaultString(target.Address, targetAddr)),
"averageTime": 0,
"packetLoss": packetLoss,
"message": message,
"fromChainType": 1,
"forwardMode": "nftables",
"nftRuleStatus": status,
"nftRuleHash": ruleHash,
"inPort": inPort,
"protocols": protocols,
"targetAddr": targetAddr,
}
return []map[string]interface{}{result}, nil
}
func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []diagnosisWorkItem, error) {
if forward == nil {
return "", nil, errForwardNotFound
@@ -1487,10 +1597,14 @@ func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int, options diag
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
pingCount := options.pingCount
if pingCount <= 0 {
pingCount = 4
}
res, err := h.sendNodeCommandWithTimeout(nodeID, "TcpPing", map[string]interface{}{
"ip": ip,
"port": port,
"count": 4,
"count": pingCount,
"timeout": options.pingTimeoutMS,
}, options.commandTimeout, false, false)
if err != nil {
@@ -1517,12 +1631,16 @@ func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int, op
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
pingCount := options.pingCount
if pingCount <= 0 {
pingCount = 4
}
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
IP: strings.TrimSpace(ip),
Port: port,
Count: 4,
Count: pingCount,
Timeout: options.pingTimeoutMS,
Protocol: "tcp",
})
@@ -1757,6 +1875,7 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
services := make([]map[string]interface{}, 0, 2)
targets := splitRemoteTargets(forward.RemoteAddr)
strategy := strings.TrimSpace(forward.Strategy)
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(forward.ProxyProtocol, forward.ProxyProtocolReceive, forward.ProxyProtocolSend)
if strategy == "" {
strategy = "fifo"
}
@@ -1801,12 +1920,16 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
if runtimeLimiters.TrafficLimiter != "" {
service["limiter"] = runtimeLimiters.TrafficLimiter
}
if forward.ProxyProtocol > 0 {
if proxyProtocolReceive > 0 {
serviceMetadata := ensureServiceMetadata(service)
serviceMetadata["proxyProtocol"] = proxyProtocolReceive
}
if proxyProtocolSend > 0 {
handlerConfig := service["handler"].(map[string]interface{})
if handlerConfig["metadata"] == nil {
handlerConfig["metadata"] = map[string]interface{}{}
}
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
handlerConfig["metadata"].(map[string]interface{})["proxyProtocol"] = proxyProtocolSend
}
if protocol == "udp" {
listenerMetadata := map[string]interface{}{
@@ -1819,10 +1942,8 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
}
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
if service["metadata"] == nil {
service["metadata"] = map[string]interface{}{}
}
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
serviceMetadata := ensureServiceMetadata(service)
serviceMetadata["interface"] = node.InterfaceName
}
services = append(services, service)
}
@@ -1841,6 +1962,25 @@ func buildForwarderNodes(targets []string) []map[string]interface{} {
return nodes
}
func ensureServiceMetadata(service map[string]interface{}) map[string]interface{} {
if service["metadata"] == nil {
service["metadata"] = map[string]interface{}{}
}
metadata, ok := service["metadata"].(map[string]interface{})
if !ok {
metadata = map[string]interface{}{}
service["metadata"] = metadata
}
return metadata
}
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
if send == 0 && legacy > 0 {
send = legacy
}
return receive, send
}
func processServerAddress(serverAddr string) string {
serverAddr = normalizeServerAddressInput(serverAddr)
if serverAddr == "" {
@@ -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)
}
}
@@ -221,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)
@@ -1014,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())
@@ -1061,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())
}
@@ -1109,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
+37 -1
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
@@ -131,7 +134,19 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) {
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
defer h.jobsWG.Done()
ticker := time.NewTicker(time.Minute)
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 {
@@ -144,6 +159,27 @@ func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
}
}
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()
+35 -5
View File
@@ -417,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 {
@@ -424,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
@@ -432,7 +439,6 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
}
now := time.Now().UnixMilli()
forwardMode := defaultNodeForwardMode(strings.TrimSpace(asString(req["forwardMode"])))
if err := h.repo.UpdateNode(id,
asString(req["name"]),
serverIP,
@@ -2143,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
@@ -2333,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
}
@@ -2492,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 {
@@ -4722,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(),
)
@@ -10,6 +10,7 @@ import (
"net/http/httptest"
"path/filepath"
"strings"
"sync"
"testing"
"time"
@@ -20,6 +21,7 @@ import (
)
type fakeNftablesManager struct {
mu sync.Mutex
testErr error
reconcileErr error
reconcileHit int
@@ -33,11 +35,15 @@ type fakeNftablesManager struct {
}
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
@@ -47,16 +53,20 @@ func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHCon
return runtimenft.ApplyResult{
NodeID: plan.NodeID,
Script: "table inet flvx {}",
Hashes: map[int64]string{plan.NodeID: "hash"},
Hashes: runtimenft.PlanHashes(plan),
}, nil
}
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
f.mu.Lock()
defer f.mu.Unlock()
f.clearHit++
return f.clearErr
}
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.collectHit++
f.lastConfig = cfg
if f.collectErr != nil {
@@ -65,6 +75,18 @@ func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.
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
@@ -158,6 +180,49 @@ func TestNodeNftablesReconcileEndpointPersistsBindings(t *testing.T) {
}
}
func TestStartBackgroundJobsReconcilesNftablesRulesAtStartup(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
seedNftablesSSHConfig(t, h, fixture.nodeID)
tunnelID := seedTunnelForNftables(t, h, "nft-startup-tunnel", fixture.nodeID)
seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
manager := &fakeNftablesManager{}
h.nftablesManager = manager
h.StartBackgroundJobs()
t.Cleanup(h.StopBackgroundJobs)
waitForCondition(t, time.Second, func() bool {
return manager.reconcileCount() > 0
}, "nftables startup reconcile")
}
func TestStartBackgroundJobsCollectsNftablesTrafficImmediatelyAndUsesFastInterval(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
manager := &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
}}
h.nftablesManager = manager
oldInterval := nftablesTrafficCollectInterval
nftablesTrafficCollectInterval = 20 * time.Millisecond
t.Cleanup(func() { nftablesTrafficCollectInterval = oldInterval })
h.StartBackgroundJobs()
t.Cleanup(h.StopBackgroundJobs)
waitForCondition(t, time.Second, func() bool {
return manager.collectCount() >= 2
}, "immediate and repeated nftables traffic collection")
}
func TestNftablesTrafficCollectIntervalDefaultsToThirtySeconds(t *testing.T) {
if nftablesTrafficCollectInterval != 30*time.Second {
t.Fatalf("expected default nftables traffic collection interval 30s, got %s", nftablesTrafficCollectInterval)
}
}
func TestNodeNftablesClearEndpointClearsBindings(t *testing.T) {
fixture := setupNftablesHandler(t)
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
@@ -280,6 +345,41 @@ func TestNodeUpdatePreservesExistingNftablesSecretsWhenFieldsOmitted(t *testing.
}
}
func TestNodeUpdateSkipsAgentProtocolCommandForNftablesNode(t *testing.T) {
fixture := setupNftablesHandler(t)
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
"id": fixture.nodeID,
"name": "nft-node-updated",
"serverIp": "198.51.100.10",
"serverIpV4": "198.51.100.10",
"port": "1000-65535",
"forwardMode": "nftables",
"http": 1,
"tls": 1,
"socks": 1,
"sshConfig": map[string]interface{}{
"host": "203.0.113.30",
"port": 22,
"username": "admin",
"authType": "password",
"sudoMode": "none",
},
})
res := httptest.NewRecorder()
fixture.handler.nodeUpdate(res, req)
assertNftablesSuccessWithBody(t, res)
cfg, err := fixture.handler.repo.GetNodeSSHConfig(fixture.nodeID)
if err != nil {
t.Fatalf("load ssh config: %v", err)
}
if !cfg.Password.Valid || cfg.Password.String != "secret" {
t.Fatalf("expected password secret to be preserved, got %+v", cfg)
}
}
func TestValidateNftablesForwardRequestRejectsHostnameTarget(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
@@ -421,6 +521,61 @@ func TestTunnelBatchRedeployUsesNftablesReconcile(t *testing.T) {
}
}
func TestDiagnoseForwardRuntimeReturnsNftablesRuleStatus(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
tunnelID := seedTunnelForNftables(t, h, "nft-diagnose-tunnel", fixture.nodeID)
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
now := time.Now().UnixMilli()
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
ForwardID: forward.ID,
NodeID: fixture.nodeID,
InPort: 20000,
Protocols: "tcp,udp",
TargetAddr: "203.0.113.9:8080",
RuleHash: "hash-a",
Status: runtimenft.StatusApplied,
}, now); err != nil {
t.Fatalf("seed nft binding: %v", err)
}
payload, err := h.diagnoseForwardRuntime(context.Background(), &forwardRecord{
ID: forward.ID,
Name: forward.Name,
TunnelID: tunnelID,
RemoteAddr: "203.0.113.9:8080",
})
if err != nil {
t.Fatalf("diagnose forward: %v", err)
}
results, ok := payload["results"].([]map[string]interface{})
if !ok || len(results) != 1 {
t.Fatalf("expected one nftables diagnosis result, got %#v", payload["results"])
}
result := results[0]
if result["forwardMode"] != "nftables" || result["nftRuleStatus"] != runtimenft.StatusApplied {
t.Fatalf("expected nftables applied result, got %#v", result)
}
if result["success"] != true {
t.Fatalf("expected nftables diagnosis success, got %#v", result)
}
if !strings.Contains(asString(result["message"]), "已下发") {
t.Fatalf("expected applied message, got %#v", result["message"])
}
}
func waitForCondition(t *testing.T, timeout time.Duration, condition func() bool, description string) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if condition() {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("timed out waiting for %s", description)
}
func setupNftablesHandler(t *testing.T) nftablesTestFixture {
t.Helper()
@@ -490,7 +645,7 @@ func seedForwardForNftables(t *testing.T, h *Handler, tunnelID, nodeID int64, re
now := time.Now().UnixMilli()
forwardID, err := h.repo.CreateForwardTx(
1, "admin", "nft-forward", tunnelID, remoteAddr, "fifo", now, 1,
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0,
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0, 0, 0,
)
if err != nil {
t.Fatalf("create forward: %v", err)
@@ -3,6 +3,7 @@ package handler
import (
"fmt"
"net"
"net/netip"
"strings"
)
@@ -75,6 +76,12 @@ func IsValidNodeAddress(addr string) error {
if strings.ContainsAny(addr, "/?") {
return fmt.Errorf("address must not contain path or query parameters")
}
// A bare IPv6 literal contains multiple colons, so net.SplitHostPort treats
// it as a malformed host:port pair. Accept IP literals before attempting
// host:port parsing; netip also handles scoped IPv6 addresses.
if _, err := netip.ParseAddr(addr); err == nil {
return nil
}
_, _, err := net.SplitHostPort(addr)
if err != nil {
@@ -0,0 +1,46 @@
package handler
import "testing"
func TestIssue515IsValidNodeAddressAcceptsBareIPv6(t *testing.T) {
for _, addr := range []string{
"2001:db8::1",
"::1",
"fe80::1%eth0",
} {
t.Run(addr, func(t *testing.T) {
if err := IsValidNodeAddress(addr); err != nil {
t.Fatalf("expected bare IPv6 address %q to be accepted: %v", addr, err)
}
})
}
}
func TestIsValidNodeAddressKeepsExistingAddressForms(t *testing.T) {
for _, addr := range []string{
"203.0.113.10",
"node.example.com",
"node.example.com:6365",
"[2001:db8::1]:6365",
} {
t.Run(addr, func(t *testing.T) {
if err := IsValidNodeAddress(addr); err != nil {
t.Fatalf("expected node address %q to be accepted: %v", addr, err)
}
})
}
}
func TestIsValidNodeAddressRejectsURLComponents(t *testing.T) {
for _, addr := range []string{
"https://node.example.com",
"node.example.com/path",
"node.example.com?transport=tcp",
} {
t.Run(addr, func(t *testing.T) {
if err := IsValidNodeAddress(addr); err == nil {
t.Fatalf("expected node address %q to be rejected", addr)
}
})
}
}
@@ -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)
}
}
}
+55 -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" }
@@ -487,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 {
@@ -593,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.
+54 -45
View File
@@ -424,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
}
@@ -944,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").
@@ -979,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
@@ -991,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
@@ -2195,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
@@ -2630,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 {
@@ -120,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 {
@@ -155,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 {
@@ -190,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 {
@@ -226,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"
@@ -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 ""
@@ -726,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
}
@@ -819,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
}
@@ -1298,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
@@ -207,20 +207,23 @@ func (r *Repository) ListActiveForwardsByNode(nodeID int64) ([]model.ForwardReco
}
rows := make([]model.ForwardRecord, 0, len(forwards))
for _, f := range forwards {
proxyProtocolReceive, proxyProtocolSend := normalizeForwardProxyProtocol(f.ProxyProtocol, f.ProxyProtocolReceive, f.ProxyProtocolSend)
rows = append(rows, model.ForwardRecord{
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
ID: f.ID,
UserID: f.UserID,
UserName: f.UserName,
Name: f.Name,
TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
MaxConn: f.MaxConn,
IPMaxConn: f.IPMaxConn,
IPSpeedID: f.IPSpeedID,
ProxyProtocol: f.ProxyProtocol,
ProxyProtocolReceive: proxyProtocolReceive,
ProxyProtocolSend: proxyProtocolSend,
})
}
for i := range rows {
@@ -237,3 +240,10 @@ func defaultString(value, fallback string) string {
}
return value
}
func normalizeForwardProxyProtocol(legacy, receive, send int) (int, int) {
if send == 0 && legacy > 0 {
send = legacy
}
return receive, send
}
@@ -1,6 +1,7 @@
package repo
import (
"path/filepath"
"strings"
"testing"
"time"
@@ -123,6 +124,114 @@ func TestNftablesNodeModeSSHConfigAndBindingPersistence(t *testing.T) {
}
}
func TestNftablesNodeSSHConfigSurvivesRepositoryReopen(t *testing.T) {
tests := []struct {
name string
cfg NftSSHConfigInput
}{
{
name: "password",
cfg: NftSSHConfigInput{
Host: "203.0.113.10",
Port: 2222,
Username: "root",
AuthType: "password",
Password: "ssh-password",
Passphrase: "key-passphrase",
SudoMode: "sudo",
},
},
{
name: "private-key",
cfg: NftSSHConfigInput{
Host: "203.0.113.11",
Port: 2223,
Username: "admin",
AuthType: "private_key",
PrivateKey: "PRIVATE-KEY-SHOULD-PERSIST",
Passphrase: "key-passphrase",
SudoMode: "sudo_su",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "nftables-ssh.sqlite")
r, err := Open(dbPath)
if err != nil {
t.Fatalf("open repo: %v", err)
}
now := time.Now().UnixMilli()
if err := r.CreateNode(
"nft-node",
"secret",
tt.cfg.Host,
nil,
nil,
"10000-20000",
nil,
nil,
nil,
nil,
nil,
0,
0,
0,
now,
1,
"[::]",
"[::]",
1,
0,
nil,
nil,
nil,
nil,
"nftables",
); err != nil {
t.Fatalf("CreateNode: %v", err)
}
nodes, err := r.ListNodes()
if err != nil {
t.Fatalf("ListNodes: %v", err)
}
nodeID := nodes[0]["id"].(int64)
if err := r.UpsertNodeSSHConfig(nodeID, tt.cfg, now); err != nil {
t.Fatalf("UpsertNodeSSHConfig: %v", err)
}
if err := r.Close(); err != nil {
t.Fatalf("close repo: %v", err)
}
reopened, err := Open(dbPath)
if err != nil {
t.Fatalf("reopen repo: %v", err)
}
defer reopened.Close()
loaded, err := reopened.GetNodeSSHConfig(nodeID)
if err != nil {
t.Fatalf("GetNodeSSHConfig after reopen: %v", err)
}
if loaded.Host != tt.cfg.Host || loaded.Port != tt.cfg.Port || loaded.Username != tt.cfg.Username || loaded.AuthType != tt.cfg.AuthType || loaded.SudoMode != tt.cfg.SudoMode {
t.Fatalf("unexpected ssh config after reopen: %+v", loaded)
}
if tt.cfg.Password != "" && (!loaded.Password.Valid || loaded.Password.String != tt.cfg.Password) {
t.Fatalf("expected password to persist after reopen, got %+v", loaded.Password)
}
if tt.cfg.PrivateKey != "" && (!loaded.PrivateKey.Valid || loaded.PrivateKey.String != tt.cfg.PrivateKey) {
t.Fatalf("expected private key to persist after reopen, got %+v", loaded.PrivateKey)
}
if tt.cfg.Passphrase != "" && (!loaded.Passphrase.Valid || loaded.Passphrase.String != tt.cfg.Passphrase) {
t.Fatalf("expected passphrase to persist after reopen, got %+v", loaded.Passphrase)
}
})
}
}
func TestUpdateNodeWithoutForwardModePreservesExistingMode(t *testing.T) {
r, err := Open(":memory:")
if err != nil {
@@ -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
+46 -6
View File
@@ -2,6 +2,8 @@ package chain
import (
"context"
"errors"
"io"
"github.com/go-gost/core/chain"
"github.com/go-gost/core/hop"
@@ -38,11 +40,12 @@ type chainNamer interface {
}
type Chain struct {
name string
hops []hop.Hop
marker selector.Marker
metadata metadata.Metadata
logger logger.Logger
name string
hops []hop.Hop
ownedHops []hop.Hop
marker selector.Marker
metadata metadata.Metadata
logger logger.Logger
}
func NewChain(name string, opts ...ChainOption) *Chain {
@@ -61,8 +64,15 @@ func NewChain(name string, opts ...ChainOption) *Chain {
}
}
func (c *Chain) AddHop(hop hop.Hop) {
func (c *Chain) AddHop(hop hop.Hop, owned ...bool) {
c.hops = append(c.hops, hop)
isOwned := true
if len(owned) > 0 {
isOwned = owned[0]
}
if isOwned {
c.ownedHops = append(c.ownedHops, hop)
}
}
// Metadata implements metadata.Metadatable interface.
@@ -112,6 +122,36 @@ func (c *Chain) Route(ctx context.Context, network, address string, opts ...chai
return rt
}
// Retire gracefully drains resources owned by a chain that has been replaced.
func (c *Chain) Retire() {
if c == nil {
return
}
for _, h := range c.ownedHops {
if retirer, ok := h.(interface{ Retire() }); ok {
retirer.Retire()
continue
}
if closer, ok := h.(io.Closer); ok {
_ = closer.Close()
}
}
}
// Close immediately releases all resources owned by the chain.
func (c *Chain) Close() error {
if c == nil {
return nil
}
var errs []error
for _, h := range c.ownedHops {
if closer, ok := h.(io.Closer); ok {
errs = append(errs, closer.Close())
}
}
return errors.Join(errs...)
}
type chainGroup struct {
chains []chain.Chainer
selector selector.Selector[chain.Chainer]
+64
View File
@@ -0,0 +1,64 @@
package chain
import (
"context"
"testing"
corechain "github.com/go-gost/core/chain"
corehop "github.com/go-gost/core/hop"
)
type lifecycleTestHop struct {
selected int
retired int
closed int
}
func (h *lifecycleTestHop) Select(context.Context, ...corehop.SelectOption) *corechain.Node {
h.selected++
return nil
}
func (h *lifecycleTestHop) Retire() {
h.retired++
}
func (h *lifecycleTestHop) Close() error {
h.closed++
return nil
}
func TestChainRoutesThroughSharedHopWithoutOwningLifecycle(t *testing.T) {
hop := &lifecycleTestHop{}
chain := NewChain("shared-hop")
chain.AddHop(hop, false)
if route := chain.Route(context.Background(), "tcp", "example.com:443"); route == nil {
t.Fatal("route is nil")
}
if hop.selected != 1 {
t.Fatalf("shared hop selected %d times, want 1", hop.selected)
}
chain.Retire()
if err := chain.Close(); err != nil {
t.Fatalf("close chain: %v", err)
}
if hop.retired != 0 || hop.closed != 0 {
t.Fatalf("shared hop lifecycle changed: retired=%d closed=%d", hop.retired, hop.closed)
}
}
func TestChainRetiresAndClosesOwnedHop(t *testing.T) {
hop := &lifecycleTestHop{}
chain := NewChain("owned-hop")
chain.AddHop(hop)
chain.Retire()
if err := chain.Close(); err != nil {
t.Fatalf("close chain: %v", err)
}
if hop.retired != 1 || hop.closed != 1 {
t.Fatalf("owned hop lifecycle: retired=%d closed=%d, want 1/1", hop.retired, hop.closed)
}
}
+32 -1
View File
@@ -2,6 +2,8 @@ package chain
import (
"context"
"errors"
"io"
"net"
"github.com/go-gost/core/chain"
@@ -102,5 +104,34 @@ func (tr *Transport) Options() *chain.TransportOptions {
func (tr *Transport) Copy() chain.Transporter {
tr2 := &Transport{}
*tr2 = *tr
return tr
return tr2
}
// Retire prevents long-lived dialer sessions owned by an obsolete chain from
// accepting new streams while allowing existing streams to drain.
func (tr *Transport) Retire() {
if tr == nil {
return
}
if retirer, ok := tr.dialer.(interface{ Retire() }); ok {
retirer.Retire()
}
if retirer, ok := tr.connector.(interface{ Retire() }); ok {
retirer.Retire()
}
}
// Close immediately releases transport-owned dialer and connector resources.
func (tr *Transport) Close() error {
if tr == nil {
return nil
}
var errs []error
if closer, ok := tr.dialer.(io.Closer); ok {
errs = append(errs, closer.Close())
}
if closer, ok := tr.connector.(io.Closer); ok {
errs = append(errs, closer.Close())
}
return errors.Join(errs...)
}
+45
View File
@@ -0,0 +1,45 @@
package chain
import (
"context"
"net"
"testing"
corechain "github.com/go-gost/core/chain"
)
type copyTestRoute struct{}
func (copyTestRoute) Dial(context.Context, string, string, ...corechain.DialOption) (net.Conn, error) {
return nil, nil
}
func (copyTestRoute) Bind(context.Context, string, string, ...corechain.BindOption) (net.Listener, error) {
return nil, nil
}
func (copyTestRoute) Nodes() []*corechain.Node {
return nil
}
func TestTransportCopyReturnsIndependentTransport(t *testing.T) {
originalRoute := copyTestRoute{}
replacementRoute := &copyTestRoute{}
original := NewTransport(nil, nil, corechain.RouteTransportOption(originalRoute))
copied, ok := original.Copy().(*Transport)
if !ok {
t.Fatalf("copy type = %T, want *Transport", original.Copy())
}
if copied == original {
t.Fatal("Copy returned the original transport")
}
copied.Options().Route = replacementRoute
if original.Options().Route != originalRoute {
t.Fatal("mutating copied transport changed original route")
}
if copied.Options().Route != replacementRoute {
t.Fatal("copied transport did not retain its independent route")
}
}
+3 -1
View File
@@ -35,16 +35,18 @@ func ParseChain(cfg *config.ChainConfig, log logger.Logger) (chain.Chainer, erro
for _, ch := range cfg.Hops {
var hop hop.Hop
var err error
owned := false
if ch.Nodes != nil || ch.Plugin != nil {
if hop, err = hop_parser.ParseHop(ch, log); err != nil {
return nil, err
}
owned = true
} else {
hop = registry.HopRegistry().Get(ch.Name)
}
if hop != nil {
c.AddHop(hop)
c.AddHop(hop, owned)
}
}
+5 -2
View File
@@ -19,20 +19,23 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session.session == nil {
if session == nil || session.session == nil {
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session.session == nil {
if session == nil || session.session == nil {
return true
}
return session.session.IsClosed()
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+38
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
kcp_util "github.com/go-gost/x/internal/util/kcp"
"github.com/go-gost/x/internal/util/sessionretire"
mdutil "github.com/go-gost/x/metadata/util"
"github.com/go-gost/x/registry"
"github.com/xtaci/kcp-go/v5"
@@ -25,6 +26,7 @@ func init() {
type kcpDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
logger logger.Logger
md metadata
options dialer.Options
@@ -64,6 +66,9 @@ func (d *kcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOp
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -171,3 +176,36 @@ func (d *kcpDialer) initSession(ctx context.Context, addr net.Addr, conn net.Pac
func (d *kcpDialer) Multiplex() bool {
return true
}
// Retire drains existing streams and closes their backing sessions once idle.
func (d *kcpDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
// Close immediately releases all cached multiplex sessions.
func (d *kcpDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *kcpDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+17
View File
@@ -0,0 +1,17 @@
package kcp
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*kcpDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil {
return true
}
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+40
View File
@@ -11,6 +11,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
"github.com/go-gost/x/registry"
)
@@ -21,6 +22,7 @@ func init() {
type mtcpDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
logger logger.Logger
md metadata
options dialer.Options
@@ -55,6 +57,9 @@ func (d *mtcpDialer) Multiplex() bool {
func (d *mtcpDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -88,6 +93,10 @@ func (d *mtcpDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
if d.md.handshakeTimeout > 0 {
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
@@ -129,3 +138,34 @@ func (d *mtcpDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
}
return &muxSession{conn: conn, session: session}, nil
}
func (d *mtcpDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mtcpDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mtcpDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mtcp
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mtcpDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil {
return true
}
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+40
View File
@@ -12,6 +12,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
"github.com/go-gost/x/registry"
)
@@ -22,6 +23,7 @@ func init() {
type mtlsDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
logger logger.Logger
md metadata
options dialer.Options
@@ -56,6 +58,9 @@ func (d *mtlsDialer) Multiplex() bool {
func (d *mtlsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -89,6 +94,10 @@ func (d *mtlsDialer) Handshake(ctx context.Context, conn net.Conn, options ...di
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
if d.md.handshakeTimeout > 0 {
conn.SetDeadline(time.Now().Add(d.md.handshakeTimeout))
@@ -136,3 +145,34 @@ func (d *mtlsDialer) initSession(ctx context.Context, conn net.Conn) (*muxSessio
}
return &muxSession{conn: conn, session: session}, nil
}
func (d *mtlsDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mtlsDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mtlsDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mtls
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mtlsDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+14
View File
@@ -20,13 +20,24 @@ func (session *muxSession) Accept() (net.Conn, error) {
}
func (session *muxSession) Close() error {
if session == nil {
return nil
}
if session.session == nil {
if session.conn != nil {
conn := session.conn
session.conn = nil
return conn.Close()
}
return nil
}
return session.session.Close()
}
func (session *muxSession) IsClosed() bool {
if session == nil {
return true
}
if session.session == nil {
return true
}
@@ -34,5 +45,8 @@ func (session *muxSession) IsClosed() bool {
}
func (session *muxSession) NumStreams() int {
if session == nil || session.session == nil {
return 0
}
return session.session.NumStreams()
}
+40
View File
@@ -12,6 +12,7 @@ import (
"github.com/go-gost/core/logger"
md "github.com/go-gost/core/metadata"
"github.com/go-gost/x/internal/util/mux"
"github.com/go-gost/x/internal/util/sessionretire"
ws_util "github.com/go-gost/x/internal/util/ws"
"github.com/go-gost/x/registry"
"github.com/gorilla/websocket"
@@ -25,6 +26,7 @@ func init() {
type mwsDialer struct {
sessions map[string]*muxSession
sessionMutex sync.Mutex
retired bool
tlsEnabled bool
md metadata
options dialer.Options
@@ -70,6 +72,9 @@ func (d *mwsDialer) Multiplex() bool {
func (d *mwsDialer) Dial(ctx context.Context, addr string, opts ...dialer.DialOption) (conn net.Conn, err error) {
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
return nil, net.ErrClosed
}
session, ok := d.sessions[addr]
if session != nil && session.IsClosed() {
@@ -108,6 +113,10 @@ func (d *mwsDialer) Handshake(ctx context.Context, conn net.Conn, options ...dia
d.sessionMutex.Lock()
defer d.sessionMutex.Unlock()
if d.retired {
conn.Close()
return nil, net.ErrClosed
}
session, ok := d.sessions[opts.Addr]
if session != nil && session.conn != conn {
@@ -208,3 +217,34 @@ func (d *mwsDialer) keepAlive(conn ws_util.WebsocketConn) {
conn.SetWriteDeadline(time.Time{})
}
}
func (d *mwsDialer) Retire() {
for _, session := range d.detachSessions() {
sessionretire.Gracefully(session)
}
}
func (d *mwsDialer) Close() error {
var errs []error
for _, session := range d.detachSessions() {
errs = append(errs, session.Close())
}
return errors.Join(errs...)
}
func (d *mwsDialer) detachSessions() []*muxSession {
if d == nil {
return nil
}
d.sessionMutex.Lock()
d.retired = true
sessions := make([]*muxSession, 0, len(d.sessions))
for _, session := range d.sessions {
if session != nil {
sessions = append(sessions, session)
}
}
d.sessions = make(map[string]*muxSession)
d.sessionMutex.Unlock()
return sessions
}
+30
View File
@@ -0,0 +1,30 @@
package mws
import (
"context"
"errors"
"net"
"testing"
)
func TestRetiredDialerRejectsNewConnections(t *testing.T) {
dialer := NewDialer().(*mwsDialer)
dialer.Retire()
if _, err := dialer.Dial(context.Background(), "127.0.0.1:1"); !errors.Is(err, net.ErrClosed) {
t.Fatalf("Dial error = %v, want net.ErrClosed", err)
}
}
func TestSessionCloseReleasesPreHandshakeConnection(t *testing.T) {
conn, peer := net.Pipe()
defer peer.Close()
session := &muxSession{conn: conn}
if err := session.Close(); err != nil {
t.Fatalf("Close: %v", err)
}
if !session.IsClosed() {
t.Fatal("pre-handshake session still reports open after Close")
}
}
+49 -8
View File
@@ -3,6 +3,7 @@ package hop
import (
"context"
"encoding/json"
"errors"
"io"
"net"
"sort"
@@ -92,6 +93,7 @@ type chainHop struct {
nodes []*chain.Node
mu sync.RWMutex
cancelFunc context.CancelFunc
stopOnce sync.Once
options options
}
@@ -383,13 +385,52 @@ func (p *chainHop) parseNode(r io.Reader) ([]*chain.Node, error) {
return nodes, nil
}
func (p *chainHop) Close() error {
p.cancelFunc()
if p.options.fileLoader != nil {
p.options.fileLoader.Close()
func (p *chainHop) stopReload() {
if p == nil {
return
}
if p.options.redisLoader != nil {
p.options.redisLoader.Close()
}
return nil
p.stopOnce.Do(func() {
p.cancelFunc()
if p.options.fileLoader != nil {
p.options.fileLoader.Close()
}
if p.options.redisLoader != nil {
p.options.redisLoader.Close()
}
if p.options.httpLoader != nil {
p.options.httpLoader.Close()
}
})
}
func (p *chainHop) Retire() {
if p == nil {
return
}
p.stopReload()
for _, node := range p.Nodes() {
if node == nil || node.Options().Transport == nil {
continue
}
if retirer, ok := node.Options().Transport.(interface{ Retire() }); ok {
retirer.Retire()
}
}
}
func (p *chainHop) Close() error {
if p == nil {
return nil
}
p.stopReload()
var errs []error
for _, node := range p.Nodes() {
if node == nil || node.Options().Transport == nil {
continue
}
if closer, ok := node.Options().Transport.(io.Closer); ok {
errs = append(errs, closer.Close())
}
}
return errors.Join(errs...)
}
@@ -0,0 +1,58 @@
package sessionretire
import "time"
const (
defaultIdleGrace = time.Second
defaultPollPeriod = 100 * time.Millisecond
)
// Session is the lifecycle surface shared by the multiplexed dialers.
type Session interface {
Close() error
IsClosed() bool
NumStreams() int
}
// Gracefully closes a retired session after all existing streams have drained.
// A short idle grace covers the Dial/Handshake hand-off used by several dialers.
func Gracefully(session Session) {
if session == nil {
return
}
go waitUntilIdle(session, defaultIdleGrace, defaultPollPeriod)
}
func waitUntilIdle(session Session, idleGrace, pollPeriod time.Duration) {
if session == nil {
return
}
if idleGrace <= 0 {
idleGrace = defaultIdleGrace
}
if pollPeriod <= 0 {
pollPeriod = defaultPollPeriod
}
ticker := time.NewTicker(pollPeriod)
defer ticker.Stop()
var idleSince time.Time
for {
if session.IsClosed() {
_ = session.Close()
return
}
if session.NumStreams() == 0 {
if idleSince.IsZero() {
idleSince = time.Now()
} else if time.Since(idleSince) >= idleGrace {
_ = session.Close()
return
}
} else {
idleSince = time.Time{}
}
<-ticker.C
}
}
@@ -0,0 +1,59 @@
package sessionretire
import (
"sync"
"testing"
"time"
)
type testSession struct {
mu sync.Mutex
streams int
closed bool
}
func (s *testSession) Close() error {
s.mu.Lock()
s.closed = true
s.mu.Unlock()
return nil
}
func (s *testSession) IsClosed() bool {
s.mu.Lock()
defer s.mu.Unlock()
return s.closed
}
func (s *testSession) NumStreams() int {
s.mu.Lock()
defer s.mu.Unlock()
return s.streams
}
func TestWaitUntilIdlePreservesActiveStreams(t *testing.T) {
session := &testSession{streams: 1}
done := make(chan struct{})
go func() {
waitUntilIdle(session, 20*time.Millisecond, time.Millisecond)
close(done)
}()
time.Sleep(30 * time.Millisecond)
if session.IsClosed() {
t.Fatal("active session was closed")
}
session.mu.Lock()
session.streams = 0
session.mu.Unlock()
select {
case <-done:
case <-time.After(250 * time.Millisecond):
t.Fatal("idle session was not closed")
}
if !session.IsClosed() {
t.Fatal("retired session did not close after becoming idle")
}
}
+7 -1
View File
@@ -28,7 +28,13 @@ func (r *chainRegistry) Register(name string, v chain.Chainer) error {
}
func (r *chainRegistry) replace(name string, v chain.Chainer) {
r.m.Store(name, v)
old, loaded := r.m.Swap(name, v)
if !loaded {
return
}
if retirer, ok := old.(interface{ Retire() }); ok {
retirer.Retire()
}
}
func (r *chainRegistry) Get(name string) chain.Chainer {
+26
View File
@@ -16,6 +16,15 @@ func (c testChainer) Route(context.Context, string, string, ...chain.RouteOption
return c.route
}
type retiringTestChainer struct {
testChainer
retired bool
}
func (c *retiringTestChainer) Retire() {
c.retired = true
}
type testRoute struct {
nodes []*chain.Node
}
@@ -49,3 +58,20 @@ func TestReplaceChainOverwritesExistingRegistration(t *testing.T) {
t.Fatalf("expected replacement chain route, got %#v", route)
}
}
func TestReplaceChainRetiresPreviousRegistration(t *testing.T) {
name := "replace_chain_retire_tdd"
ChainRegistry().Unregister(name)
defer ChainRegistry().Unregister(name)
old := &retiringTestChainer{}
if err := ChainRegistry().Register(name, old); err != nil {
t.Fatalf("register old chain: %v", err)
}
if err := ReplaceChain(name, testChainer{}); err != nil {
t.Fatalf("replace chain: %v", err)
}
if !old.retired {
t.Fatal("previous chain was not retired")
}
}
+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{
+289 -42
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,16 +51,51 @@ get_architecture() {
# 安装目录
INSTALL_DIR="/etc/flux_agent"
FLUX_AGENT_SYSTEMD_SERVICE_FILE="/etc/systemd/system/flux_agent.service"
FLUX_AGENT_OPENRC_SERVICE_FILE="/etc/init.d/flux_agent"
LEGACY_GOST_BINARY="/usr/local/bin/gost"
LEGACY_GOST_CONFIG_DIR="/etc/gost"
LEGACY_GOST_SERVICE_FILE_ETC="/etc/systemd/system/gost.service"
LEGACY_GOST_SERVICE_FILE_LIB="/lib/systemd/system/gost.service"
LEGACY_GOST_SERVICE_FILE_USR_LIB="/usr/lib/systemd/system/gost.service"
SERVICE_MANAGER="${SERVICE_MANAGER:-}"
# 镜像加速配置(可由面板传入或交互式询问)
PROXY_ENABLED="${PROXY_ENABLED:-}"
PROXY_URL="${PROXY_URL:-}"
ensure_alpine_runtime_dependencies() {
[[ -f /etc/alpine-release ]] || return 0
local missing_packages=()
local privileged_command=""
command -v curl >/dev/null 2>&1 || missing_packages+=(curl)
[[ -f /etc/ssl/certs/ca-certificates.crt ]] || missing_packages+=(ca-certificates)
if [[ ${#missing_packages[@]} -eq 0 ]]; then
return 0
fi
if [[ $EUID -ne 0 ]]; then
if command -v sudo >/dev/null 2>&1; then
privileged_command="sudo"
elif command -v doas >/dev/null 2>&1; then
privileged_command="doas"
else
echo "❌ Alpine 安装需要 root 权限,或已配置 sudo/doas 来安装依赖: ${missing_packages[*]}。" >&2
return 1
fi
fi
echo "📦 Alpine 缺少运行依赖,正在安装: ${missing_packages[*]}"
if [[ -n "$privileged_command" ]]; then
"$privileged_command" apk add --no-cache "${missing_packages[@]}"
else
apk add --no-cache "${missing_packages[@]}"
fi
}
# 镜像加速
maybe_proxy_url() {
local url="$1"
@@ -139,6 +201,8 @@ build_download_url() {
}
ensure_download_url_initialized() {
ensure_alpine_runtime_dependencies || return 1
if [[ -n "${DOWNLOAD_URL:-}" ]]; then
return 0
fi
@@ -256,6 +320,204 @@ write_flux_agent_config() {
"$(json_escape "$SECRET")" > "$path"
}
ensure_service_manager() {
if [[ -n "$SERVICE_MANAGER" ]]; then
case "$SERVICE_MANAGER" in
systemd|openrc)
return 0
;;
*)
echo "❌ 不支持的服务管理器: $SERVICE_MANAGER" >&2
return 1
;;
esac
fi
# Alpine uses OpenRC even if a systemctl compatibility command happens to be installed.
if [[ -f /etc/alpine-release ]]; then
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
SERVICE_MANAGER="openrc"
return 0
fi
elif command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
SERVICE_MANAGER="systemd"
return 0
elif command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
SERVICE_MANAGER="openrc"
return 0
fi
echo "❌ 未检测到受支持的服务管理器(systemd 或 OpenRC)。" >&2
return 1
}
flux_agent_service_exists() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
[[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" ]] || \
systemctl list-units --full -all 2>/dev/null | grep -Fq "flux_agent.service"
;;
openrc)
[[ -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]
;;
esac
}
stop_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl stop flux_agent 2>/dev/null || true
;;
openrc)
rc-service flux_agent stop 2>/dev/null || true
;;
esac
}
disable_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl disable flux_agent 2>/dev/null || true
;;
openrc)
rc-update del flux_agent default 2>/dev/null || true
;;
esac
}
write_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
mkdir -p "$(dirname "$FLUX_AGENT_SYSTEMD_SERVICE_FILE")"
cat > "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" <<EOF
[Unit]
Description=Flux_agent Proxy Service
After=network.target
[Service]
WorkingDirectory=$INSTALL_DIR
ExecStart=$INSTALL_DIR/flux_agent
Restart=on-failure
StandardOutput=null
StandardError=null
[Install]
WantedBy=multi-user.target
EOF
;;
openrc)
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
cat > "$FLUX_AGENT_OPENRC_SERVICE_FILE" <<EOF
#!/sbin/openrc-run
name="flux_agent"
description="Flux_agent Proxy Service"
command="$INSTALL_DIR/flux_agent"
directory="$INSTALL_DIR"
command_background="yes"
pidfile="/run/\${RC_SVCNAME}.pid"
output_log="/dev/null"
error_log="/dev/null"
depend() {
need net
}
EOF
chmod +x "$FLUX_AGENT_OPENRC_SERVICE_FILE"
;;
esac
}
enable_and_start_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl daemon-reload
systemctl enable flux_agent
systemctl start flux_agent
;;
openrc)
rc-update add flux_agent default
rc-service flux_agent start
;;
esac
}
start_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl start flux_agent
;;
openrc)
rc-service flux_agent start
;;
esac
}
flux_agent_service_is_active() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl is-active --quiet flux_agent
;;
openrc)
rc-service flux_agent status >/dev/null 2>&1
;;
esac
}
flux_agent_service_status() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl is-active flux_agent 2>/dev/null || true
;;
openrc)
rc-service flux_agent status 2>/dev/null || true
;;
esac
}
remove_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
rm -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE"
systemctl daemon-reload 2>/dev/null || true
;;
openrc)
rm -f "$FLUX_AGENT_OPENRC_SERVICE_FILE"
;;
esac
}
flux_agent_service_status_hint() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
echo "systemctl status flux_agent --no-pager"
;;
openrc)
echo "rc-service flux_agent status"
;;
esac
}
cleanup_legacy_gost_installation() {
local matched_service_files=()
local service_file=""
@@ -278,7 +540,8 @@ cleanup_legacy_gost_installation() {
return 0
fi
if systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
if [[ "$SERVICE_MANAGER" == "systemd" ]] && \
systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
systemctl stop gost 2>/dev/null || true
systemctl disable gost 2>/dev/null || true
fi
@@ -297,7 +560,7 @@ cleanup_legacy_gost_installation() {
rm -f "$LEGACY_GOST_CONFIG_DIR/gost"
fi
if [[ "$removed_service_file" == "1" ]]; then
if [[ "$removed_service_file" == "1" && "$SERVICE_MANAGER" == "systemd" ]]; then
systemctl daemon-reload 2>/dev/null || true
fi
}
@@ -341,19 +604,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 +657,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 +691,7 @@ update_flux_agent() {
# 检查并安装 tcpkill
check_and_install_tcpkill
ensure_service_manager || return 1
# 先下载新版本
echo "⬇️ 下载最新版本..."
@@ -455,9 +704,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 +718,7 @@ update_flux_agent() {
# 重启服务
echo "🔄 重启服务..."
systemctl start flux_agent
start_flux_agent_service
echo "✅ 更新完成,服务已重新启动。"
}
@@ -477,6 +726,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 +735,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 +753,6 @@ uninstall_flux_agent() {
echo "🧹 删除安装目录: $INSTALL_DIR"
fi
# 重载 systemd
systemctl daemon-reload
echo "✅ 卸载完成"
}
+152 -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"
@@ -283,6 +399,7 @@ test_cleanup_legacy_gost_installation_removes_service_and_binary() (
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
SERVICE_MANAGER="systemd"
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
[Unit]
Description=Gost Proxy Service
@@ -326,6 +443,7 @@ test_cleanup_legacy_gost_installation_preserves_unrelated_gost() (
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
SERVICE_MANAGER="systemd"
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<'EOF'
[Unit]
Description=Unrelated Gost Service
@@ -353,6 +471,35 @@ EOF
[[ "$systemctl_calls" != *"disable gost"* ]] || fail "cleanup_legacy_gost_installation should not disable unrelated gost services"
)
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
LEGACY_GOST_SERVICE_FILE_ETC=$(mktemp)
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
SERVICE_MANAGER="openrc"
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
[Unit]
WorkingDirectory=$LEGACY_GOST_CONFIG_DIR
ExecStart=$LEGACY_GOST_CONFIG_DIR/gost
EOF
: > "$LEGACY_GOST_CONFIG_DIR/config.json"
: > "$LEGACY_GOST_CONFIG_DIR/gost.json"
local systemctl_calls=""
systemctl() {
systemctl_calls+=$'\n'"$*"
return 1
}
cleanup_legacy_gost_installation >/dev/null
[[ -z "$systemctl_calls" ]] || fail "OpenRC cleanup should not invoke systemctl"
[[ ! -e "$LEGACY_GOST_SERVICE_FILE_ETC" ]] || fail "OpenRC cleanup should still remove the legacy service file"
)
test_install_script_accepts_proxy_url_env_without_prompt() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
@@ -504,8 +651,12 @@ test_update_flux_agent_skips_proxy_prompt_when_not_installed
test_install_flux_agent_preserves_legacy_gost_when_download_fails
test_update_flux_agent_preserves_legacy_gost_when_download_fails
test_install_flux_agent_writes_json_safe_config
test_install_script_bootstraps_bash_for_alpine
test_install_flux_agent_uses_openrc
test_remove_flux_agent_service_uses_openrc
test_cleanup_legacy_gost_installation_removes_service_and_binary
test_cleanup_legacy_gost_installation_preserves_unrelated_gost
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc
test_install_script_accepts_proxy_url_env_without_prompt
test_panel_install_script_can_disable_proxy
test_panel_install_script_recomputes_compose_urls_after_prompt
@@ -514,4 +665,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"
+2
View File
@@ -217,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) =>
+18
View File
@@ -82,6 +82,8 @@ export interface ForwardApiItem {
ipSpeedLimitName?: string;
maxConn?: number;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
inx?: number;
[key: string]: unknown;
}
@@ -422,6 +424,8 @@ export interface ForwardMutationPayload {
ipSpeedId?: number | null;
maxConn?: number;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
}
export interface SpeedLimitMutationPayload {
@@ -591,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;
+8 -5
View File
@@ -28,12 +28,13 @@ function DialogClose({
return <DialogPrimitive.Close data-slot="dialog-close" {...props} />;
}
function DialogOverlay({
className,
...props
}: React.ComponentProps<typeof DialogPrimitive.Overlay>) {
const DialogOverlay = React.forwardRef<
React.ElementRef<typeof DialogPrimitive.Overlay>,
React.ComponentPropsWithoutRef<typeof DialogPrimitive.Overlay>
>(({ className, ...props }, ref) => {
return (
<DialogPrimitive.Overlay
ref={ref}
className={cn(
"fixed inset-0 z-50 bg-black/30 backdrop-blur-md data-[state=open]:animate-in data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=open]:fade-in-0",
className,
@@ -42,7 +43,9 @@ function DialogOverlay({
{...props}
/>
);
}
});
DialogOverlay.displayName = DialogPrimitive.Overlay.displayName;
function DialogContent({
className,
@@ -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 (
+248 -23
View File
@@ -69,6 +69,7 @@ import {
getNodeList,
pauseForwardService,
resumeForwardService,
resetForwardFlow,
diagnoseForward,
updateForwardOrder,
getConfigByName,
@@ -130,6 +131,8 @@ interface Forward {
ipSpeedId?: number | null;
ipSpeedLimitName?: string;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
}
interface Tunnel {
@@ -169,6 +172,8 @@ interface ForwardForm {
ipSpeedId: number | null;
maxConn?: number;
proxyProtocol?: number;
proxyProtocolReceive?: number;
proxyProtocolSend?: number;
}
interface ForwardUserGroup {
@@ -226,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 => {
@@ -600,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,
}));
};
@@ -750,6 +765,7 @@ const SortableTableRow = ({
handleEdit,
handleDelete,
handleDiagnose,
handleResetFlow,
showAddressModal,
formatFlow,
}: any) => {
@@ -905,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"
@@ -944,6 +983,7 @@ const SortableCompactTableRow = ({
handleEdit,
handleDelete,
handleDiagnose,
handleResetFlow,
showAddressModal,
hasMultipleAddresses,
formatFlow,
@@ -1130,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"
@@ -1275,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] =
@@ -1337,6 +1405,8 @@ export default function ForwardPage() {
ipSpeedId: null,
maxConn: 0,
proxyProtocol: 0,
proxyProtocolReceive: 0,
proxyProtocolSend: 0,
});
const [inIpTouched, setInIpTouched] = useState(false);
@@ -2128,6 +2198,8 @@ export default function ForwardPage() {
ipMaxConn: 0,
ipSpeedId: null,
proxyProtocol: 0,
proxyProtocolReceive: 0,
proxyProtocolSend: 0,
});
setErrors({});
setModalOpen(true);
@@ -2152,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);
@@ -2163,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;
@@ -2285,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);
@@ -2301,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);
@@ -3933,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"
@@ -3976,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"
@@ -4278,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>
@@ -4297,6 +4442,7 @@ export default function ForwardPage() {
handleDelete={handleDelete}
handleDiagnose={handleDiagnose}
handleEdit={handleEdit}
handleResetFlow={handleResetFlow}
handleServiceToggle={handleServiceToggle}
hasMultipleAddresses={hasMultipleAddresses}
selectMode={selectMode}
@@ -4592,6 +4738,7 @@ export default function ForwardPage() {
handleDelete={handleDelete}
handleDiagnose={handleDiagnose}
handleEdit={handleEdit}
handleResetFlow={handleResetFlow}
handleServiceToggle={
handleServiceToggle
}
@@ -4968,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="规则限速"
@@ -5123,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={{
@@ -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,249 @@ 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 {};
};
type TunnelTopologyHop = TunnelQualityHopApiItem & {
errorMessage?: string;
};
interface TunnelTopologyPath {
key: string;
hops: TunnelTopologyHop[];
alternativeNodeIndex?: number;
}
const tunnelTopologyHopKey = (fromNodeId: number, toNodeId: number) =>
`${fromNodeId}:${toNodeId}`;
const buildTunnelTopologyPaths = (
details: TunnelQualityChainDetailsApiItem,
): TunnelTopologyPath[] => {
const primaryHops = details.primaryPath ?? [];
const candidates = details.candidateHops ?? [];
if (primaryHops.length === 0) {
const publicCandidates = candidates.filter(
(candidate) => candidate.toRole === "target",
);
const selected = publicCandidates.find((candidate) => candidate.selected);
const paths: TunnelTopologyPath[] = [];
if (selected) {
paths.push({ key: "primary-public", hops: [selected] });
}
for (const candidate of publicCandidates) {
if (candidate.selected) continue;
paths.push({
key: `alternative-public-${candidate.fromNodeId}`,
hops: [candidate],
alternativeNodeIndex: 0,
});
}
return paths;
}
const primaryNodeIds = [
primaryHops[0].fromNodeId,
...primaryHops.map((hop) => hop.toNodeId),
];
const internalCandidates = candidates.filter(
(candidate) => candidate.toRole !== "target",
);
const candidateHopMap = new Map<string, TunnelQualityCandidateHopApiItem>();
const alternativeNodes = new Map<
string,
{ column: number; nodeId: number }
>();
for (const candidate of internalCandidates) {
candidateHopMap.set(
tunnelTopologyHopKey(candidate.fromNodeId, candidate.toNodeId),
candidate,
);
const sourceColumn = candidate.hopIndex;
const targetColumn = candidate.hopIndex + 1;
if (
sourceColumn >= 0 &&
sourceColumn < primaryNodeIds.length &&
candidate.fromNodeId !== primaryNodeIds[sourceColumn]
) {
alternativeNodes.set(`${sourceColumn}:${candidate.fromNodeId}`, {
column: sourceColumn,
nodeId: candidate.fromNodeId,
});
}
if (
targetColumn >= 0 &&
targetColumn < primaryNodeIds.length &&
candidate.toNodeId !== primaryNodeIds[targetColumn]
) {
alternativeNodes.set(`${targetColumn}:${candidate.toNodeId}`, {
column: targetColumn,
nodeId: candidate.toNodeId,
});
}
}
const paths: TunnelTopologyPath[] = [{ key: "primary", hops: primaryHops }];
for (const alternative of alternativeNodes.values()) {
const nodeIds = [...primaryNodeIds];
nodeIds[alternative.column] = alternative.nodeId;
const hops: TunnelTopologyHop[] = [];
for (let index = 0; index < nodeIds.length - 1; index += 1) {
const fromNodeId = nodeIds[index];
const toNodeId = nodeIds[index + 1];
const usesPrimaryEdge =
fromNodeId === primaryNodeIds[index] &&
toNodeId === primaryNodeIds[index + 1];
const hop = usesPrimaryEdge
? primaryHops[index]
: candidateHopMap.get(tunnelTopologyHopKey(fromNodeId, toNodeId));
if (!hop) break;
hops.push(hop);
}
if (hops.length === primaryHops.length) {
paths.push({
key: `alternative-${alternative.column}-${alternative.nodeId}`,
hops,
alternativeNodeIndex: alternative.column,
});
}
}
return paths;
};
function TunnelTopologyPathRow({ path }: { path: TunnelTopologyPath }) {
return (
<div className="flex min-w-max items-center py-2">
{path.hops.map((hop, index) => {
const hasError =
Boolean(hop.errorMessage) || hop.latency < 0 || hop.loss > 0;
const colorClass =
hop.latency < 0 || hop.errorMessage
? "text-danger"
: hop.loss > 0
? "text-warning"
: "text-success";
const borderColor = hasError ? "border-danger" : "";
return (
<React.Fragment
key={`${path.key}-${hop.fromNodeId}-${hop.toNodeId}-${index}`}
>
{index === 0 ? (
<TopologyNodeChip
isAlternative={path.alternativeNodeIndex === 0}
name={hop.fromNodeName}
/>
) : null}
<div
className="relative mx-1 flex min-w-[70px] shrink-0 flex-col items-center justify-center"
title={hop.errorMessage}
>
<span
className={`mb-1 text-[10px] font-mono leading-none ${colorClass}`}
>
{hop.latency >= 0 && !hop.errorMessage
? `${hop.latency.toFixed(0)}ms`
: "超时"}
</span>
<div
className={`relative flex h-[2px] w-full items-center justify-end bg-default-200 ${hop.latency < 0 || hop.errorMessage ? "!bg-danger" : ""}`}
>
<ArrowRight
className={`absolute -right-2 z-10 h-3.5 w-3.5 rounded-full bg-background p-[1px] ${colorClass}`}
/>
</div>
<span
className={`mt-1.5 text-[10px] font-mono leading-none ${hop.errorMessage || hop.loss > 0 ? "text-warning" : "text-default-400"}`}
>
{hop.errorMessage ? "探测失败" : `${hop.loss.toFixed(0)}% 丢包`}
</span>
</div>
<TopologyNodeChip
borderColor={borderColor}
isAlternative={path.alternativeNodeIndex === index + 1}
name={hop.toNodeName}
/>
</React.Fragment>
);
})}
</div>
);
}
function TopologyNodeChip({
name,
isAlternative = false,
borderColor = "",
}: {
name: string;
isAlternative?: boolean;
borderColor?: string;
}) {
return (
<Chip
className={`shrink-0 font-mono shadow-sm ${borderColor}`}
size="sm"
variant="flat"
>
<span className="flex items-center gap-1.5">
<span>{name}</span>
{isAlternative ? (
<span className="rounded-full bg-warning/20 px-1.5 py-0.5 text-[9px] font-semibold leading-none text-warning">
备选
</span>
) : null}
</span>
</Chip>
);
}
const ForwardingChainTopology = React.memo(function ForwardingChainTopology({
hopsStr,
}: {
hopsStr?: string;
}) {
const details = useMemo(
() => parseTunnelQualityChainDetails(hopsStr),
[hopsStr],
);
const topologyPaths = useMemo(
() => buildTunnelTopologyPaths(details),
[details],
);
if (topologyPaths.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,62 +723,24 @@ 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" : "";
return (
<React.Fragment key={index}>
{index === 0 && (
<Chip
className="shrink-0 font-mono shadow-sm"
size="sm"
variant="flat"
>
{hop.fromNodeName}
</Chip>
)}
<div className="flex flex-col items-center justify-center min-w-[70px] mx-1 shrink-0 relative">
<span
className={`text-[10px] font-mono leading-none mb-1 ${colorClass}`}
>
{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 className="max-h-80 space-y-1 overflow-auto pb-1">
{topologyPaths.map((path, index) => (
<div
key={path.key}
className={
index === 0
? "overflow-x-auto"
: "overflow-x-auto border-t border-dashed border-divider/60"
}
>
<TunnelTopologyPathRow path={path} />
</div>
))}
</div>
</CardBody>
</Card>
);
}
});
export function TunnelMonitorView({
viewMode = "grid",
@@ -562,6 +762,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 +821,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 +854,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 +866,12 @@ export function TunnelMonitorView({
setQualityLoading(false);
}
} else {
void loadMonitorTunnelQualityEnabled();
void loadTunnelQualityConfig();
}
if (typeof intervalSec === "number") {
setTunnelQualityIntervalSec(
parseTunnelQualityIntervalSeconds(String(intervalSec)),
);
}
};
@@ -673,7 +886,7 @@ export function TunnelMonitorView({
handleMonitorTunnelQualityEnabledChanged as EventListener,
);
};
}, [loadMonitorTunnelQualityEnabled]);
}, [loadTunnelQualityConfig]);
useEffect(() => {
if (tunnels.length > 0 && !initialHistoryFetched.current) {
@@ -726,7 +939,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 +1001,7 @@ export function TunnelMonitorView({
qualityTimerRef.current = window.setInterval(() => {
void loadQuality({ silent: true });
}, QUALITY_POLL_INTERVAL);
}, tunnelQualityIntervalSec * 1000);
return () => {
if (qualityTimerRef.current) {
@@ -796,7 +1009,7 @@ export function TunnelMonitorView({
qualityTimerRef.current = null;
}
};
}, [loadQuality, monitorTunnelQualityEnabled]);
}, [loadQuality, monitorTunnelQualityEnabled, tunnelQualityIntervalSec]);
// --- Load quality history for detail chart ---
const loadQualityHistory = useCallback(
@@ -1065,7 +1278,7 @@ export function TunnelMonitorView({
{monitorTunnelQualityEnabled ? (
<>
<LiveDot />
<span>自动探测中(每秒测试,30秒上报)</span>
<span>{`自动探测中(${tunnelQualityIntervalLabel(tunnelQualityIntervalSec)}测试)`}</span>
</>
) : (
<>
@@ -1129,7 +1342,7 @@ export function TunnelMonitorView({
{monitorTunnelQualityEnabled ? (
<>
<LiveDot />
<span>每秒探测 · 更新于 {lastQualityUpdate}</span>
<span>{`${tunnelQualityIntervalLabel(tunnelQualityIntervalSec)}探测 · 更新于 ${lastQualityUpdate}`}</span>
</>
) : (
<>
@@ -49,6 +49,41 @@ function useModalContext() {
return React.useContext(ModalContext);
}
interface ScrollPosition {
element: HTMLElement | null;
left: number;
top: number;
}
function captureScrollPositions(): ScrollPosition[] {
const positions: ScrollPosition[] = [
{ element: null, left: window.scrollX, top: window.scrollY },
];
for (const element of Array.from(
document.querySelectorAll<HTMLElement>("main, [data-scroll-container]"),
)) {
positions.push({
element,
left: element.scrollLeft,
top: element.scrollTop,
});
}
return positions;
}
function restoreScrollPositions(positions: ScrollPosition[]) {
for (const position of positions) {
if (position.element) {
position.element.scrollLeft = position.left;
position.element.scrollTop = position.top;
} else {
window.scrollTo(position.left, position.top);
}
}
}
type ModalSize = "sm" | "md" | "lg" | "xl" | "2xl" | "4xl" | "full";
function mapSize(size: ModalSize | undefined) {
@@ -97,6 +132,46 @@ export function Modal({
scrollBehavior,
size,
}: ModalProps) {
const previousScrollPositionsRef = React.useRef<ScrollPosition[] | null>(
null,
);
// Radix focus management and scroll locking can move an ancestor scroll
// container when a modal is opened from a card/grid item. Capture the
// current positions before the open render and restore them after focus
// settles so opening a modal never changes the page position.
React.useLayoutEffect(() => {
return () => {
if (!isOpen) {
previousScrollPositionsRef.current = captureScrollPositions();
}
};
}, [isOpen]);
React.useLayoutEffect(() => {
const positions = previousScrollPositionsRef.current;
if (!isOpen || !positions) {
return;
}
restoreScrollPositions(positions);
let nestedFrame = 0;
const frame = window.requestAnimationFrame(() => {
restoreScrollPositions(positions);
nestedFrame = window.requestAnimationFrame(() =>
restoreScrollPositions(positions),
);
});
previousScrollPositionsRef.current = null;
return () => {
window.cancelAnimationFrame(frame);
window.cancelAnimationFrame(nestedFrame);
};
}, [isOpen]);
const handleOpenChange = (open: boolean) => {
onOpenChange?.(open);
if (!open) {