Compare commits

...

20 Commits

Author SHA1 Message Date
sagit 2269f2e2d5 fix: include pnpm config in frontend image build (#552) 2026-09-01 15:48:54 +08:00
sagit da7bef88f1 fix: harden panel updates and nftables deployment (#551)
Fix panel update deployment discovery and rollback safety, add nftables compatibility and atomic replacement, and restore reproducible frontend CI installs.
2026-09-01 15:37:33 +08:00
sagit f26014579b fix: stabilize federated agent runtime updates (#550) 2026-08-25 17:04:04 +08:00
sagit 5041c722c9 fix: harden license lifecycle (#547) 2026-08-13 09:25:05 +08:00
sagit c56798e991 fix: accept existing license machine binding (#546) 2026-08-12 14:12:46 +08:00
sagit 40e96f3592 fix: preserve active per-IP traffic limiters (#545) 2026-08-12 11:31:03 +08:00
sagit 0e24b53a5b fix: preserve license state across panel upgrades (#544) 2026-08-11 17:04:20 +08:00
sagit a820c49c94 fix(agent): harden Alpine OpenRC installation (#543)
Ensure Alpine installs use OpenRC without invoking systemd cleanup paths, and install missing download dependencies.
2026-08-08 11:45:39 +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
88 changed files with 6408 additions and 453 deletions
+1 -1
View File
@@ -22,7 +22,7 @@ jobs:
node-version: '20.19.0'
- name: Install pnpm
run: npm install -g pnpm
run: npm install -g pnpm@10.28.1
- name: Install dependencies
run: pnpm install --frozen-lockfile
+1 -1
View File
@@ -229,6 +229,7 @@ jobs:
docker buildx build \
--platform linux/amd64,linux/arm64 \
--build-arg KEYGEN_ACCOUNT_ID=${{ secrets.KEYGEN_ACCOUNT_ID }} \
--push \
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:latest \
-t ${{ env.REGISTRY }}/${OWNER}/flux-panel-backend:${VERSION} \
@@ -431,4 +432,3 @@ jobs:
gh release upload "${VERSION}" ./artifacts/gost-arm64.sha256 --clobber
echo "✅ GOST 二进制文件更新完成"
+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`
不需要数据库迁移或新增依赖。
+2 -1
View File
@@ -7,7 +7,8 @@ RUN go mod download
COPY . .
ARG TARGETOS
ARG TARGETARCH
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -o /out/paneld ./cmd/paneld
ARG KEYGEN_ACCOUNT_ID
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} env ${TARGETARCH:+GOARCH=${TARGETARCH}} go build -ldflags="-X 'go-backend/internal/license.AccountID=${KEYGEN_ACCOUNT_ID}'" -o /out/paneld ./cmd/paneld
FROM docker:27-cli AS dockercli
@@ -29,6 +29,10 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if value, gated := h.unlicensedPublicBrandValue(configName); gated {
response.WriteJSON(w, response.OK(map[string]string{"name": configName, "value": value}))
return
}
cfg, err := h.repo.GetConfigByName(configName)
if err != nil {
@@ -42,3 +46,22 @@ func (h *Handler) getPublicConfigByName(w http.ResponseWriter, r *http.Request)
response.WriteJSON(w, response.OK(cfg))
}
func (h *Handler) unlicensedPublicBrandValue(configName string) (string, bool) {
switch configName {
case "app_name", "app_logo", "app_favicon", "hide_footer_brand":
default:
return "", false
}
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if isCommercial == "true" {
return "", false
}
if configName == "app_name" {
return "FLVX", true
}
if configName == "hide_footer_brand" {
return "false", true
}
return "", true
}
@@ -30,6 +30,40 @@ func TestPublicConfigGetAllowsBrandKeys(t *testing.T) {
assertHandlerCode(t, resp, 0)
}
func TestPublicBrandConfigFallsBackWithoutCommercialLicense(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
seedConfigValue(t, r, "app_name", "Paid Brand")
seedConfigValue(t, r, "app_logo", "logo-data")
seedConfigValue(t, r, "app_favicon", "favicon-data")
seedConfigValue(t, r, "hide_footer_brand", "true")
seedConfigValue(t, r, "is_commercial", "false")
for name, want := range map[string]string{
"app_name": "FLVX",
"app_logo": "",
"app_favicon": "",
"hide_footer_brand": "false",
} {
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertHandlerConfigValue(t, resp, name, want)
}
}
func TestPublicBrandConfigUsesSavedValuesWithCommercialLicense(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
seedConfigValue(t, r, "app_name", "Paid Brand")
seedConfigValue(t, r, "is_commercial", "true")
req := httptest.NewRequest(http.MethodPost, "/api/v1/public/config/get", bytes.NewBufferString(`{"name":"app_name"}`))
req.Header.Set("Content-Type", "application/json")
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertHandlerConfigValue(t, resp, "app_name", "Paid Brand")
}
func TestPublicConfigGetRejectsSensitiveKeys(t *testing.T) {
router, _ := setupConfigAccessTestRouter(t)
@@ -82,6 +116,56 @@ func TestConfigGetAllowsSensitiveKeysForAdmin(t *testing.T) {
assertHandlerConfigValue(t, resp, "jwt_secret", "jwt-secret")
}
func TestConfigGetNeverReturnsLicenseCredentials(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
seedConfigValue(t, r, "license_key", "license-secret")
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
for _, name := range []string{"license_key", "license_machine_id", "machine_fingerprint"} {
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/get", bytes.NewBufferString(`{"name":"`+name+`"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", adminToken)
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
assertHandlerCodeMsg(t, resp, 403, "禁止访问系统授权凭据")
}
}
func TestConfigListNeverReturnsLicenseCredentials(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
seedConfigValue(t, r, "license_key", "license-secret")
seedConfigValue(t, r, "license_machine_id", "machine-id")
seedConfigValue(t, r, "machine_fingerprint", "fingerprint-secret")
seedConfigValue(t, r, "is_commercial", "true")
req := httptest.NewRequest(http.MethodPost, "/api/v1/config/list", nil)
req.Header.Set("Authorization", adminToken)
resp := httptest.NewRecorder()
router.ServeHTTP(resp, req)
var out struct {
Code int `json:"code"`
Data map[string]string `json:"data"`
}
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
t.Fatalf("decode response: %v", err)
}
if out.Code != 0 || out.Data["is_commercial"] != "true" {
t.Fatalf("unexpected config response: %+v", out)
}
if _, ok := out.Data["license_key"]; ok {
t.Fatal("license_key must not be returned")
}
if _, ok := out.Data["license_machine_id"]; ok {
t.Fatal("license_machine_id must not be returned")
}
if _, ok := out.Data["machine_fingerprint"]; ok {
t.Fatal("machine_fingerprint must not be returned")
}
}
func TestConfigUpdateAllowsSensitiveKeysForAdmin(t *testing.T) {
router, _ := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
@@ -154,7 +238,7 @@ func TestConfigUpdateSingleAllowsCloudflareSecretKeyWrite(t *testing.T) {
}
}
func TestConfigUpdateAllowsLicenseKeyWrite(t *testing.T) {
func TestConfigUpdateRejectsLicenseKeyWrite(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
@@ -165,18 +249,13 @@ func TestConfigUpdateAllowsLicenseKeyWrite(t *testing.T) {
router.ServeHTTP(resp, req)
assertHandlerCode(t, resp, 0)
cfg, err := r.GetConfigByName("license_key")
if err != nil {
t.Fatalf("get config: %v", err)
}
if cfg == nil || cfg.Value != "license-secret" {
t.Fatalf("expected license_key to be updated, got %#v", cfg)
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
}
}
func TestConfigUpdateSingleAllowsLicenseKeyWrite(t *testing.T) {
func TestConfigUpdateSingleRejectsLicenseKeyWrite(t *testing.T) {
router, r := setupConfigAccessTestRouter(t)
adminToken := mustGenerateConfigAccessToken(t, 1, "admin_user", 0)
@@ -187,14 +266,9 @@ func TestConfigUpdateSingleAllowsLicenseKeyWrite(t *testing.T) {
router.ServeHTTP(resp, req)
assertHandlerCode(t, resp, 0)
cfg, err := r.GetConfigByName("license_key")
if err != nil {
t.Fatalf("get config: %v", err)
}
if cfg == nil || cfg.Value != "license-secret" {
t.Fatalf("expected license_key to be updated, got %#v", cfg)
assertHandlerCodeMsg(t, resp, -1, "该配置由系统管理")
if cfg, err := r.GetConfigByName("license_key"); err != nil || cfg != nil {
t.Fatalf("license_key should not be written, got %#v err=%v", cfg, err)
}
}
@@ -14,6 +14,7 @@ import (
"time"
"go-backend/internal/http/client"
runtimenft "go-backend/internal/runtime/nftables"
"go-backend/internal/store/model"
"go-backend/internal/ws"
)
@@ -58,6 +59,7 @@ type diagnosisWorkItem struct {
type diagnosisExecOptions struct {
commandTimeout time.Duration
pingTimeoutMS int
pingCount int
timeoutMessage string
}
@@ -772,6 +774,9 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
if ctx == nil {
ctx = context.Background()
}
if payload, handled, err := h.diagnoseNftablesForwardRuntime(forward); handled || err != nil {
return payload, err
}
forwardName, workItems, err := h.prepareForwardDiagnosis(forward)
if err != nil {
return nil, err
@@ -787,6 +792,111 @@ func (h *Handler) diagnoseForwardRuntime(ctx context.Context, forward *forwardRe
return payload, nil
}
func (h *Handler) diagnoseNftablesForwardRuntime(forward *forwardRecord) (map[string]interface{}, bool, error) {
if forward == nil {
return nil, false, errForwardNotFound
}
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
if err != nil {
return nil, false, err
}
if !nftMode {
return nil, false, nil
}
if len(entryNodeIDs) == 0 {
return nil, true, errors.New("nftables 转发缺少入口节点")
}
targets, err := resolveDiagnosisTargets(forward.RemoteAddr)
if err != nil {
return nil, true, err
}
results, err := h.buildNftablesForwardDiagnosisResults(forward, entryNodeIDs[0], targets)
if err != nil {
return nil, true, err
}
payload := map[string]interface{}{
"forwardName": forward.Name,
"timestamp": time.Now().UnixMilli(),
"results": results,
}
return payload, true, nil
}
func (h *Handler) buildNftablesForwardDiagnosisResults(forward *forwardRecord, nodeID int64, targets []diagnosisTarget) ([]map[string]interface{}, error) {
if h == nil || h.repo == nil {
return nil, errors.New("handler not initialized")
}
node, err := h.getNodeRecord(nodeID)
if err != nil {
return nil, err
}
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
if err != nil {
return nil, err
}
var binding *model.NftRuleBinding
for i := range bindings {
if bindings[i].ForwardID == forward.ID {
binding = &bindings[i]
break
}
}
target := diagnosisTarget{}
if len(targets) > 0 {
target = targets[0]
}
status := "missing"
message := "nftables 规则未下发"
success := false
inPort := 0
protocols := ""
targetAddr := strings.TrimSpace(forward.RemoteAddr)
ruleHash := ""
if binding != nil {
status = strings.ToLower(strings.TrimSpace(binding.Status))
inPort = binding.InPort
protocols = strings.TrimSpace(binding.Protocols)
targetAddr = strings.TrimSpace(binding.TargetAddr)
ruleHash = strings.TrimSpace(binding.RuleHash)
if status == "" {
status = "pending"
}
if status == runtimenft.StatusApplied {
success = true
message = "nftables 规则已下发"
} else if strings.TrimSpace(binding.LastError) != "" {
message = binding.LastError
} else {
message = "nftables 规则未完成下发"
}
}
packetLoss := 100
if success {
packetLoss = 0
}
result := map[string]interface{}{
"success": success,
"nodeName": node.Name,
"nodeId": strconv.FormatInt(nodeID, 10),
"targetIp": target.IP,
"targetPort": target.Port,
"description": fmt.Sprintf("nftables规则(%s)->目标(%s)", node.Name, defaultString(target.Address, targetAddr)),
"averageTime": 0,
"packetLoss": packetLoss,
"message": message,
"fromChainType": 1,
"forwardMode": "nftables",
"nftRuleStatus": status,
"nftRuleHash": ruleHash,
"inPort": inPort,
"protocols": protocols,
"targetAddr": targetAddr,
}
return []map[string]interface{}{result}, nil
}
func (h *Handler) prepareForwardDiagnosis(forward *forwardRecord) (string, []diagnosisWorkItem, error) {
if forward == nil {
return "", nil, errForwardNotFound
@@ -1487,10 +1597,14 @@ func (h *Handler) tcpPingViaNode(nodeID int64, ip string, port int, options diag
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
pingCount := options.pingCount
if pingCount <= 0 {
pingCount = 4
}
res, err := h.sendNodeCommandWithTimeout(nodeID, "TcpPing", map[string]interface{}{
"ip": ip,
"port": port,
"count": 4,
"count": pingCount,
"timeout": options.pingTimeoutMS,
}, options.commandTimeout, false, false)
if err != nil {
@@ -1517,12 +1631,16 @@ func (h *Handler) tcpPingViaRemoteNode(node *nodeRecord, ip string, port int, op
if options.pingTimeoutMS <= 0 {
options.pingTimeoutMS = int(diagnosisCommandTimeout / time.Millisecond)
}
pingCount := options.pingCount
if pingCount <= 0 {
pingCount = 4
}
fc := client.NewFederationClientWithTimeout(options.commandTimeout)
return fc.Diagnose(remoteURL, remoteToken, h.federationLocalDomain(), client.RuntimeDiagnoseRequest{
IP: strings.TrimSpace(ip),
Port: port,
Count: 4,
Count: pingCount,
Timeout: options.pingTimeoutMS,
Protocol: "tcp",
})
@@ -659,9 +659,21 @@ func (h *Handler) sendDeleteOrphanedForwardService(nodeID int64, serviceName str
}
func (h *Handler) speedLimiterExists(name string) bool {
name = strings.TrimSpace(name)
if name == "" {
return false
}
const forwardRulePrefix = "rule_traffic_limit_"
if strings.HasPrefix(name, forwardRulePrefix) {
forwardID, err := strconv.ParseInt(strings.TrimPrefix(name, forwardRulePrefix), 10, 64)
if err != nil || forwardID <= 0 {
return false
}
forward, err := h.getForwardRecord(forwardID)
return err == nil && forward != nil && forward.IPSpeedID.Valid && forward.IPSpeedID.Int64 > 0
}
id, err := strconv.ParseInt(name, 10, 64)
if err != nil || id <= 0 {
return false
@@ -0,0 +1,38 @@
package handler
import (
"path/filepath"
"testing"
"go-backend/internal/store/repo"
)
func TestSpeedLimiterExistsPreservesForwardRuleLimiter(t *testing.T) {
r, err := repo.Open(filepath.Join(t.TempDir(), "panel.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
if err := r.DB().Exec(`
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx, ip_speed_id)
VALUES(8, 1, 'user', 'forward', 1, '127.0.0.1:80', 'fifo', 0, 0, 1, 1, 1, 0, 3),
(9, 1, 'user', 'forward-without-ip-limit', 1, '127.0.0.1:81', 'fifo', 0, 0, 1, 1, 1, 0, NULL)
`).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
h := &Handler{repo: r}
if !h.speedLimiterExists("rule_traffic_limit_8") {
t.Fatal("expected runtime limiter for existing forward to be preserved")
}
if h.speedLimiterExists("rule_traffic_limit_9") {
t.Fatal("expected runtime limiter for forward without per-IP speed limit to be treated as orphaned")
}
if h.speedLimiterExists("rule_traffic_limit_10") {
t.Fatal("expected runtime limiter for missing forward to be treated as orphaned")
}
if h.speedLimiterExists("rule_traffic_limit_invalid") {
t.Fatal("expected malformed runtime limiter name to be treated as orphaned")
}
}
@@ -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)
}
}
+63 -47
View File
@@ -5,6 +5,7 @@ import (
"database/sql"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"log"
@@ -20,7 +21,6 @@ import (
"go-backend/internal/health"
"go-backend/internal/http/middleware"
"go-backend/internal/http/response"
"go-backend/internal/license"
"go-backend/internal/metrics"
"go-backend/internal/monitoring"
runtimenft "go-backend/internal/runtime/nftables"
@@ -42,10 +42,12 @@ type Handler struct {
captchaMu sync.Mutex
captchaTokens map[string]int64
jobsMu sync.Mutex
jobsCancel context.CancelFunc
jobsStarted bool
jobsWG sync.WaitGroup
jobsMu sync.Mutex
jobsCancel context.CancelFunc
jobsStarted bool
jobsWG sync.WaitGroup
fingerprintMu sync.Mutex
licenseValidationMu sync.Mutex
upgradeMu sync.Mutex
systemUpgradeMu sync.Mutex
@@ -221,6 +223,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/forward/force-delete", h.forwardForceDelete)
mux.HandleFunc("/api/v1/forward/pause", h.forwardPause)
mux.HandleFunc("/api/v1/forward/resume", h.forwardResume)
mux.HandleFunc("/api/v1/forward/reset-flow", h.forwardResetFlow)
mux.HandleFunc("/api/v1/forward/diagnose", h.forwardDiagnose)
mux.HandleFunc("/api/v1/forward/diagnose/stream", h.forwardDiagnoseStream)
mux.HandleFunc("/api/v1/forward/update-order", h.forwardUpdateOrder)
@@ -399,6 +402,10 @@ func (h *Handler) getConfigByName(w http.ResponseWriter, r *http.Request) {
return
}
configName := strings.ToLower(strings.TrimSpace(req.Name))
if configName == "license_key" || configName == "license_machine_id" || configName == "machine_fingerprint" {
response.WriteJSON(w, response.Err(403, "禁止访问系统授权凭据"))
return
}
if repo.IsSensitiveConfigKey(configName) && !isAdminRequest(r) {
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
@@ -434,11 +441,13 @@ func (h *Handler) getConfigs(w http.ResponseWriter, r *http.Request) {
return
}
ctxClaims := r.Context().Value(middleware.ClaimsContextKey)
if claims, ok := ctxClaims.(auth.Claims); !ok || claims.RoleID != 0 {
delete(cfgMap, "license_key")
delete(cfgMap, "cloudflare_secret_key")
delete(cfgMap, "jwt_secret")
claims, isAdmin := ctxClaims.(auth.Claims)
if !isAdmin || claims.RoleID != 0 {
cfgMap = repo.FilterSensitiveConfigs(cfgMap)
}
delete(cfgMap, "license_key")
delete(cfgMap, "license_machine_id")
delete(cfgMap, "machine_fingerprint")
response.WriteJSON(w, response.OK(cfgMap))
}
@@ -877,10 +886,16 @@ func (h *Handler) flowUpload(w http.ResponseWriter, r *http.Request) {
}
func (h *Handler) getOrCreateMachineFingerprint() (string, error) {
fp, _ := h.repo.GetViteConfigValue("machine_fingerprint")
h.fingerprintMu.Lock()
defer h.fingerprintMu.Unlock()
fp, err := h.repo.GetViteConfigValue("machine_fingerprint")
if fp != "" {
return fp, nil
}
if err != nil && !errors.Is(err, sql.ErrNoRows) {
return "", err
}
newFp := uuid.New().String()
now := time.Now().UnixMilli()
@@ -907,56 +922,35 @@ func (h *Handler) licenseActivate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("授权码不能为空"))
return
}
h.licenseValidationMu.Lock()
defer h.licenseValidationMu.Unlock()
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
fingerprint, err := h.getOrCreateMachineFingerprint()
valResp, err := h.validateLicenseForMachine(key)
if err != nil {
response.WriteJSON(w, response.ErrDefault("生成设备指纹失败"))
return
}
client := license.NewKeygenClient(accountID, "")
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
if err != nil {
response.WriteJSON(w, response.ErrDefault("连接授权服务器失败: "+err.Error()))
log.Printf("license activation failed: %v", err)
response.WriteJSON(w, response.ErrDefault(licenseValidationErrorMessage(err)))
return
}
if !valResp.Meta.Valid {
if valResp.Meta.Code == "NO_MACHINES" || valResp.Meta.Code == "NO_MACHINE" || valResp.Meta.Code == "MACHINE_SCOPE_REQUIRED" || valResp.Meta.Code == "FINGERPRINT_SCOPE_MISMATCH" {
// Needs machine activation
client.Token = key
err = client.ActivateMachine(valResp.Data.ID, fingerprint)
if err != nil {
// Translate specific error messages or log them
response.WriteJSON(w, response.ErrDefault("设备绑定失败: "+err.Error()))
return
}
// Validation might still fail with scope if we don't query via machine id, but since activate machine succeeded
// we can consider the license valid for our simple usecase
} else {
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
return
}
response.WriteJSON(w, response.ErrDefault("授权码无效或已过期 (Code: "+valResp.Meta.Code+")"))
return
}
now := time.Now().UnixMilli()
if err := h.repo.UpsertConfig("license_key", key, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.repo.UpsertConfig("is_commercial", "true", now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
expiry := valResp.Data.Attributes.Expiry
if expiry == "" {
expiry = "never"
}
if err := h.repo.UpsertConfig("license_expiry", expiry, now); err != nil {
licenseState := map[string]string{
"license_key": key,
"is_commercial": "true",
"license_expiry": expiry,
}
if valResp.MachineID != "" {
licenseState["license_machine_id"] = valResp.MachineID
}
if err := h.repo.UpsertConfigs(licenseState, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -998,6 +992,10 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if repo.IsSystemManagedConfigKey(key) {
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
return
}
if protectedKeys[key] && isCommercial != "true" {
response.WriteJSON(w, response.ErrDefault("需要商业版授权"))
@@ -1014,6 +1012,7 @@ func (h *Handler) updateConfigs(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
h.notifyTunnelQualityConfigChanged(key)
}
response.WriteJSON(w, response.OKEmpty())
@@ -1039,6 +1038,10 @@ func (h *Handler) updateSingleConfig(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
return
}
if repo.IsSystemManagedConfigKey(name) {
response.WriteJSON(w, response.ErrDefault("该配置由系统管理"))
return
}
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if (name == "app_name" || name == "app_logo" || name == "app_favicon" || name == "hide_footer_brand") && isCommercial != "true" {
@@ -1061,6 +1064,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 +1113,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
+57 -12
View File
@@ -2,11 +2,12 @@ package handler
import (
"context"
"log"
"time"
"go-backend/internal/license"
)
var nftablesTrafficCollectInterval = 30 * time.Second
func (h *Handler) StartBackgroundJobs() {
if h == nil || h.repo == nil {
return
@@ -35,6 +36,7 @@ func (h *Handler) StartBackgroundJobs() {
func (h *Handler) runValidateLicenseJob(ctx context.Context) {
defer h.jobsWG.Done()
h.validateLicenseJob()
ticker := time.NewTicker(12 * time.Hour)
defer ticker.Stop()
@@ -52,22 +54,25 @@ func (h *Handler) validateLicenseJob() {
if h == nil || h.repo == nil {
return
}
accountID := "1bc96cac-09de-4cf4-af34-26afdad63a90"
h.licenseValidationMu.Lock()
defer h.licenseValidationMu.Unlock()
key, _ := h.repo.GetViteConfigValue("license_key")
isCommercial, _ := h.repo.GetViteConfigValue("is_commercial")
if key == "" || isCommercial != "true" {
if key == "" {
return // Nothing to validate
}
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
client := license.NewKeygenClient(accountID, "")
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
valResp, err := h.validateLicenseForMachine(key)
if err != nil {
// Network error or timeout. Grace period by not revoking immediately here.
// Network and decode failures have no validation response, so retain the
// current state as a grace period. A rejected machine binding still has
// the original invalid response and must not stay commercially enabled.
if licenseValidationErrorIsDefinitive(valResp, err) {
now := time.Now().UnixMilli()
_ = h.repo.UpsertConfig("is_commercial", "false", now)
}
return
}
@@ -81,7 +86,14 @@ func (h *Handler) validateLicenseJob() {
if expiry == "" {
expiry = "never"
}
_ = h.repo.UpsertConfig("license_expiry", expiry, now)
licenseState := map[string]string{
"is_commercial": "true",
"license_expiry": expiry,
}
if valResp.MachineID != "" {
licenseState["license_machine_id"] = valResp.MachineID
}
_ = h.repo.UpsertConfigs(licenseState, now)
}
}
@@ -131,7 +143,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 +168,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()
@@ -0,0 +1,109 @@
package handler
import (
"errors"
"fmt"
"net/http"
"os"
"strings"
"go-backend/internal/license"
)
var newLicenseClient = license.NewKeygenClient
func keygenAccountID() string {
if value := strings.TrimSpace(license.AccountID); value != "" {
return value
}
return strings.TrimSpace(os.Getenv("KEYGEN_ACCOUNT_ID"))
}
func licenseNeedsMachineActivation(code string) bool {
switch strings.ToUpper(strings.TrimSpace(code)) {
case "NO_MACHINES", "NO_MACHINE", "MACHINE_SCOPE_REQUIRED", "FINGERPRINT_SCOPE_MISMATCH":
return true
default:
return false
}
}
func licenseValidationErrorIsDefinitive(validation *license.ValidateResponse, err error) bool {
if err == nil {
return validation != nil && !validation.Meta.Valid
}
var apiErr *license.APIError
if !errors.As(err, &apiErr) {
return false
}
if apiErr.StatusCode == http.StatusTooManyRequests || apiErr.StatusCode >= http.StatusInternalServerError {
return false
}
return apiErr.Operation == "activate machine" && validation != nil && !validation.Meta.Valid
}
func licenseValidationErrorMessage(err error) string {
if strings.Contains(err.Error(), "keygen account id is not configured") {
return "授权服务配置错误"
}
var apiErr *license.APIError
if errors.As(err, &apiErr) {
if apiErr.HasCode("MACHINE_LIMIT_EXCEEDED") {
return "授权设备数量已达上限"
}
if apiErr.StatusCode == http.StatusUnauthorized || apiErr.StatusCode == http.StatusForbidden {
return "授权码无效或无权绑定设备"
}
}
return "连接授权服务器失败,请稍后重试"
}
func (h *Handler) validateLicenseForMachine(key string) (*license.ValidateResponse, error) {
fingerprint, err := h.getOrCreateMachineFingerprint()
if err != nil {
return nil, fmt.Errorf("prepare machine fingerprint: %w", err)
}
storedMachineID, _ := h.repo.GetViteConfigValue("license_machine_id")
accountID := keygenAccountID()
if accountID == "" {
return nil, fmt.Errorf("keygen account id is not configured")
}
client := newLicenseClient(accountID, "")
var validation *license.ValidateResponse
if storedMachineID != "" {
validation, err = client.ValidateKeyWithMachine(key, fingerprint, storedMachineID)
} else {
validation, err = client.ValidateKeyWithFingerprint(key, fingerprint)
}
if err != nil {
return nil, err
}
if validation.Meta.Valid || !licenseNeedsMachineActivation(validation.Meta.Code) {
if validation.Meta.Valid {
validation.MachineID = storedMachineID
}
return validation, nil
}
client.Token = key
machineID, err := client.ActivateMachine(validation.Data.ID, fingerprint)
if err != nil {
return validation, err
}
if machineID == "" {
machineID, err = client.GetMachineID(fingerprint)
if err != nil {
return validation, fmt.Errorf("retrieve activated machine: %w", err)
}
}
validation, err = client.ValidateKeyWithMachine(key, fingerprint, machineID)
if err != nil {
return nil, err
}
if validation.Meta.Valid {
validation.MachineID = machineID
}
return validation, nil
}
@@ -0,0 +1,358 @@
package handler
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"time"
"go-backend/internal/license"
"go-backend/internal/store/repo"
)
func TestValidateLicenseJobRepairsMissingMachineBinding(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
var validations atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
if validations.Add(1) == 1 {
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
return
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
case strings.Contains(req.URL.Path, "/machines/"):
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
fingerprint, err := r.GetViteConfigValue("machine_fingerprint")
if err != nil || strings.TrimSpace(fingerprint) == "" {
t.Fatalf("expected persisted machine fingerprint, got value=%q err=%v", fingerprint, err)
}
if got := validations.Load(); got != 2 {
t.Fatalf("validation calls = %d, want 2", got)
}
}
func TestValidateLicenseJobAcceptsExistingMachineActivation(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
var validations atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
if validations.Add(1) == 1 {
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
return
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
case strings.Contains(req.URL.Path, "/machines/"):
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "never")
if got := validations.Load(); got != 2 {
t.Fatalf("validation calls = %d, want 2", got)
}
}
func TestLicenseActivateRequiresSuccessfulPostActivationValidation(t *testing.T) {
r := openLicenseTestRepository(t)
var validations atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
code := "NO_MACHINE"
if validations.Add(1) > 1 {
code = "FINGERPRINT_SCOPE_MISMATCH"
}
_, _ = fmt.Fprintf(w, `{"meta":{"valid":false,"code":%q},"data":{"id":"license-id","attributes":{}}}`, code)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
case strings.Contains(req.URL.Path, "/machines/"):
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
res := httptest.NewRecorder()
h.licenseActivate(res, req)
if !strings.Contains(res.Body.String(), "FINGERPRINT_SCOPE_MISMATCH") {
t.Fatalf("expected post-activation validation failure, got %s", res.Body.String())
}
assertLicenseConfig(t, r, "is_commercial", "false")
for _, name := range []string{"license_key", "license_expiry"} {
if value, err := r.GetViteConfigValue(name); err == nil || value != "" {
t.Fatalf("%s should not be persisted, got value=%q err=%v", name, value, err)
}
}
}
func TestLicenseActivatePersistsValidatedState(t *testing.T) {
r := openLicenseTestRepository(t)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"2030-01-02T00:00:00.000Z"}}}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", bytes.NewBufferString(`{"license_key":"license-secret"}`))
res := httptest.NewRecorder()
h.licenseActivate(res, req)
if !strings.Contains(res.Body.String(), `"code":0`) {
t.Fatalf("expected activation success, got %s", res.Body.String())
}
assertLicenseConfig(t, r, "license_key", "license-secret")
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
}
func TestValidateLicenseJobDowngradesWhenMachineBindingIsRejected(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"NO_MACHINE"},"data":{"id":"license-id","attributes":{}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "false")
}
func TestValidateLicenseJobRestoresCommercialStateWhenLicenseRecovers(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "false", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
if strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
return
}
http.NotFound(w, req)
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_expiry", "never")
}
func TestValidateLicenseJobUsesStoredMachineScope(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
seedLicenseConfig(t, r, "license_machine_id", "machine-id", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
if !strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key") {
t.Fatalf("unexpected request %s %s", req.Method, req.URL.Path)
}
var body struct {
Meta struct {
Scope map[string]string `json:"scope"`
} `json:"meta"`
}
if err := json.NewDecoder(req.Body).Decode(&body); err != nil {
t.Fatalf("decode request: %v", err)
}
if body.Meta.Scope["machine"] != "machine-id" || body.Meta.Scope["fingerprint"] != "fingerprint" {
t.Fatalf("unexpected validation scope: %+v", body.Meta.Scope)
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{"expiry":"never"}}}`)
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
assertLicenseConfig(t, r, "license_machine_id", "machine-id")
}
func TestValidateLicenseJobKeepsStateOnMachineLookupServerFailure(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
case strings.Contains(req.URL.Path, "/machines/"):
w.WriteHeader(http.StatusServiceUnavailable)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"SERVICE_UNAVAILABLE"}]}`)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
}
func TestValidateLicenseJobKeepsStateOnMachineLookupNotFound(t *testing.T) {
r := openLicenseTestRepository(t)
now := time.Now().UnixMilli()
seedLicenseConfig(t, r, "license_key", "license-secret", now)
seedLicenseConfig(t, r, "is_commercial", "true", now)
seedLicenseConfig(t, r, "machine_fingerprint", "fingerprint", now)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
switch {
case strings.HasSuffix(req.URL.Path, "/licenses/actions/validate-key"):
_, _ = fmt.Fprint(w, `{"meta":{"valid":false,"code":"MACHINE_SCOPE_REQUIRED"},"data":{"id":"license-id","attributes":{}}}`)
case strings.HasSuffix(req.URL.Path, "/machines"):
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"}]}`)
case strings.Contains(req.URL.Path, "/machines/"):
http.NotFound(w, req)
default:
http.NotFound(w, req)
}
}))
defer server.Close()
restoreLicenseClientFactory(t, server.URL)
h := &Handler{repo: r}
h.validateLicenseJob()
assertLicenseConfig(t, r, "is_commercial", "true")
}
func TestLicenseValidationErrorMessageDoesNotExposeKeygenResponse(t *testing.T) {
err := &license.APIError{
Operation: "activate machine",
StatusCode: http.StatusUnprocessableEntity,
Body: `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED","detail":"private detail"}]}`,
}
message := licenseValidationErrorMessage(err)
if message != "授权设备数量已达上限" || strings.Contains(message, "private detail") {
t.Fatalf("unexpected public error message %q", message)
}
}
func openLicenseTestRepository(t *testing.T) *repo.Repository {
t.Helper()
r, err := repo.Open(filepath.Join(t.TempDir(), "license.db"))
if err != nil {
t.Fatalf("repo.Open() error = %v", err)
}
t.Cleanup(func() { _ = r.Close() })
return r
}
func seedLicenseConfig(t *testing.T, r *repo.Repository, name, value string, now int64) {
t.Helper()
if err := r.UpsertConfig(name, value, now); err != nil {
t.Fatalf("UpsertConfig(%q) error = %v", name, err)
}
}
func assertLicenseConfig(t *testing.T, r *repo.Repository, name, want string) {
t.Helper()
got, err := r.GetViteConfigValue(name)
if err != nil {
t.Fatalf("GetViteConfigValue(%q) error = %v", name, err)
}
if got != want {
t.Fatalf("config %q = %q, want %q", name, got, want)
}
}
func restoreLicenseClientFactory(t *testing.T, baseURL string) {
t.Helper()
t.Setenv("KEYGEN_ACCOUNT_ID", "account-id")
previous := newLicenseClient
newLicenseClient = func(accountID, token string) *license.KeygenClient {
client := license.NewKeygenClient(accountID, token)
client.BaseURL = baseURL
return client
}
t.Cleanup(func() { newLicenseClient = previous })
}
+28 -2
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,
@@ -2496,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 {
@@ -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()
@@ -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)
}
}
}
@@ -182,6 +182,8 @@ func requiresAdmin(path string) bool {
return true
case "/api/v1/config/update", "/api/v1/config/update-single":
return true
case "/api/v1/license/activate":
return true
case "/api/v1/announcement/update":
return true
default:
@@ -197,6 +197,39 @@ func TestShouldSkipBypassesPublicConfigGet(t *testing.T) {
}
}
func TestLicenseActivateRequiresAdmin(t *testing.T) {
if !requiresAdmin("/api/v1/license/activate") {
t.Fatal("expected license activation to require admin")
}
}
func TestJWTRejectsNonAdminLicenseActivation(t *testing.T) {
secret := "unit-test-secret"
token, err := auth.GenerateToken(2, "regular_user", 1, secret)
if err != nil {
t.Fatalf("generate token: %v", err)
}
claims, err := auth.ParseClaims(token, secret)
if err != nil {
t.Fatalf("parse claims: %v", err)
}
wrapped := JWT(AuthOptions{
JWTSecret: secret,
GetUserAuthState: func(userID int64) (*auth.UserAuthState, error) {
return &auth.UserAuthState{ID: userID, RoleID: 1, Status: 1, PasswordChangedAt: claims.IatMs - 1}, nil
},
})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.OK("pass"))
}))
req := httptest.NewRequest(http.MethodPost, "/api/v1/license/activate", nil)
req.Header.Set("Authorization", token)
res := httptest.NewRecorder()
wrapped.ServeHTTP(res, req)
assertCodeMsg(t, res, 403, "权限不足,仅管理员可操作")
}
func TestJWTExpiresAfterSevenDays(t *testing.T) {
secret := "unit-test-secret"
token, err := auth.GenerateToken(1, "admin_user", 0, secret)
+127 -21
View File
@@ -6,24 +6,53 @@ import (
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
)
var AccountID string
type KeygenClient struct {
AccountID string
Token string
BaseURL string
HTTPClient *http.Client
}
type APIError struct {
Operation string
StatusCode int
Body string
}
func (e *APIError) Error() string {
return fmt.Sprintf("keygen %s failed: status %d, response: %s", e.Operation, e.StatusCode, e.Body)
}
func (e *APIError) HasCode(code string) bool {
return e != nil && hasKeygenErrorCode([]byte(e.Body), code)
}
const defaultAPIBaseURL = "https://api.keygen.sh/v1"
func NewKeygenClient(accountID, token string) *KeygenClient {
return &KeygenClient{
AccountID: accountID,
Token: token,
AccountID: accountID,
Token: token,
BaseURL: defaultAPIBaseURL,
HTTPClient: &http.Client{Timeout: 10 * time.Second},
}
}
func (c *KeygenClient) apiURL(path string) string {
baseURL := strings.TrimRight(c.BaseURL, "/")
if baseURL == "" {
baseURL = defaultAPIBaseURL
}
return fmt.Sprintf("%s/accounts/%s/%s", baseURL, c.AccountID, strings.TrimLeft(path, "/"))
}
type ValidateResponse struct {
Meta struct {
Valid bool `json:"valid"`
@@ -35,6 +64,7 @@ type ValidateResponse struct {
Expiry string `json:"expiry"`
} `json:"attributes"`
} `json:"data"`
MachineID string `json:"-"`
}
type ActivateMachineRequest struct {
@@ -54,8 +84,37 @@ type ActivateMachineRequest struct {
} `json:"data"`
}
type keygenErrorResponse struct {
Errors []struct {
Code string `json:"code"`
} `json:"errors"`
}
type MachineResponse struct {
Data struct {
ID string `json:"id"`
} `json:"data"`
}
func hasKeygenErrorCode(body []byte, code string) bool {
var resp keygenErrorResponse
if err := json.Unmarshal(body, &resp); err != nil {
return false
}
for _, item := range resp.Errors {
if strings.EqualFold(strings.TrimSpace(item.Code), code) {
return true
}
}
return false
}
func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string) (*ValidateResponse, error) {
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
return c.ValidateKeyWithMachine(key, fingerprint, "")
}
func (c *KeygenClient) ValidateKeyWithMachine(key, fingerprint, machineID string) (*ValidateResponse, error) {
url := c.apiURL("licenses/actions/validate-key")
meta := map[string]interface{}{
"key": key,
@@ -66,6 +125,14 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
"fingerprint": fingerprint,
}
}
if machineID != "" {
scope, _ := meta["scope"].(map[string]interface{})
if scope == nil {
scope = make(map[string]interface{})
meta["scope"] = scope
}
scope["machine"] = machineID
}
reqBody := map[string]interface{}{
"meta": meta,
@@ -91,7 +158,8 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
body, _ := io.ReadAll(resp.Body)
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
}
var valResp ValidateResponse
@@ -102,8 +170,44 @@ func (c *KeygenClient) ValidateKeyWithFingerprint(key string, fingerprint string
return &valResp, nil
}
func (c *KeygenClient) GetMachineID(fingerprint string) (string, error) {
machineURL := c.apiURL("machines/" + url.PathEscape(fingerprint))
req, err := http.NewRequest(http.MethodGet, machineURL, nil)
if err != nil {
return "", err
}
req.Header.Set("Accept", "application/vnd.api+json")
if c.Token != "" {
if !strings.HasPrefix(c.Token, "Bearer ") && !strings.HasPrefix(c.Token, "License ") {
req.Header.Set("Authorization", "License "+c.Token)
} else {
req.Header.Set("Authorization", c.Token)
}
}
resp, err := c.HTTPClient.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
return "", &APIError{Operation: "retrieve machine", StatusCode: resp.StatusCode, Body: string(body)}
}
var machineResp MachineResponse
if err := json.NewDecoder(resp.Body).Decode(&machineResp); err != nil {
return "", err
}
machineID := strings.TrimSpace(machineResp.Data.ID)
if machineID == "" {
return "", fmt.Errorf("failed to retrieve machine: empty machine id")
}
return machineID, nil
}
func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/licenses/actions/validate-key", c.AccountID)
url := c.apiURL("licenses/actions/validate-key")
reqBody := map[string]interface{}{
"meta": map[string]string{
@@ -130,7 +234,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("keygen api error: status %d", resp.StatusCode)
body, _ := io.ReadAll(resp.Body)
return nil, &APIError{Operation: "validate license", StatusCode: resp.StatusCode, Body: string(body)}
}
var valResp ValidateResponse
@@ -141,8 +246,8 @@ func (c *KeygenClient) ValidateKey(key string) (*ValidateResponse, error) {
return &valResp, nil
}
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
url := fmt.Sprintf("https://api.keygen.sh/v1/accounts/%s/machines", c.AccountID)
func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) (string, error) {
url := c.apiURL("machines")
var reqBody ActivateMachineRequest
reqBody.Data.Type = "machines"
@@ -165,23 +270,24 @@ func (c *KeygenClient) ActivateMachine(licenseID, fingerprint string) error {
resp, err := c.HTTPClient.Do(req)
if err != nil {
return err
return "", err
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode == http.StatusCreated || resp.StatusCode == http.StatusOK {
return nil
}
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode == http.StatusConflict || resp.StatusCode == http.StatusUnprocessableEntity {
if strings.Contains(string(body), "FINGERPRINT_TAKEN") || strings.Contains(string(body), "MACHINE_LIMIT_EXCEEDED") {
// Machine already registered to this license or limit reached because it's already us.
// The subsequent ValidateKey check will determine if the existing machine is actually us.
return nil
var machineResp MachineResponse
if json.Unmarshal(body, &machineResp) == nil {
return strings.TrimSpace(machineResp.Data.ID), nil
}
return "", nil
}
return fmt.Errorf("failed to activate machine: status %d, response: %s", resp.StatusCode, string(body))
}
if resp.StatusCode == http.StatusUnprocessableEntity && hasKeygenErrorCode(body, "FINGERPRINT_TAKEN") {
// Machine activation is idempotent. Keygen scopes fingerprint uniqueness
// to the target license, so this means the same machine is already bound.
return "", nil
}
return "", &APIError{Operation: "activate machine", StatusCode: resp.StatusCode, Body: string(body)}
}
+104
View File
@@ -0,0 +1,104 @@
package license
import (
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestValidateKeyWithMachineSendsFingerprintAndMachineScopes(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var body struct {
Meta struct {
Scope map[string]string `json:"scope"`
} `json:"meta"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode request: %v", err)
}
if body.Meta.Scope["fingerprint"] != "fingerprint" {
t.Fatalf("fingerprint scope = %q", body.Meta.Scope["fingerprint"])
}
if body.Meta.Scope["machine"] != "machine-id" {
t.Fatalf("machine scope = %q", body.Meta.Scope["machine"])
}
_, _ = fmt.Fprint(w, `{"meta":{"valid":true,"code":"VALID"},"data":{"id":"license-id","attributes":{}}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "")
client.BaseURL = server.URL
validation, err := client.ValidateKeyWithMachine("license-key", "fingerprint", "machine-id")
if err != nil || !validation.Meta.Valid {
t.Fatalf("ValidateKeyWithMachine() validation=%+v err=%v", validation, err)
}
}
func TestGetMachineIDRetrievesMachineByFingerprint(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet || !strings.HasSuffix(r.URL.Path, "/machines/fingerprint") {
t.Fatalf("unexpected request %s %s", r.Method, r.URL.Path)
}
if r.Header.Get("Authorization") != "License license-key" {
t.Fatalf("authorization = %q", r.Header.Get("Authorization"))
}
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
machineID, err := client.GetMachineID("fingerprint")
if err != nil || machineID != "machine-id" {
t.Fatalf("GetMachineID() = %q, %v", machineID, err)
}
}
func TestActivateMachineTreatsFingerprintTakenAsIdempotent(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"FINGERPRINT_TAKEN"},{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
if machineID, err := client.ActivateMachine("license-id", "fingerprint"); err != nil || machineID != "" {
t.Fatalf("ActivateMachine() error = %v, want idempotent success", err)
}
}
func TestActivateMachineReturnsCreatedMachineID(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusCreated)
_, _ = fmt.Fprint(w, `{"data":{"type":"machines","id":"machine-id"}}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
machineID, err := client.ActivateMachine("license-id", "fingerprint")
if err != nil || machineID != "machine-id" {
t.Fatalf("ActivateMachine() = %q, %v", machineID, err)
}
}
func TestActivateMachineRejectsMachineLimitWithoutFingerprintTaken(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnprocessableEntity)
_, _ = fmt.Fprint(w, `{"errors":[{"code":"MACHINE_LIMIT_EXCEEDED"}]}`)
}))
defer server.Close()
client := NewKeygenClient("account-id", "license-key")
client.BaseURL = server.URL
_, err := client.ActivateMachine("license-id", "fingerprint")
if err == nil || !strings.Contains(err.Error(), "MACHINE_LIMIT_EXCEEDED") {
t.Fatalf("ActivateMachine() error = %v, want machine limit failure", err)
}
}
@@ -0,0 +1,52 @@
package monitoring
import (
"fmt"
"strconv"
"strings"
)
const (
ConfigTunnelQualityProbeIntervalSec = "monitor_tunnel_quality_interval_sec"
DefaultTunnelQualityProbeIntervalSec = 1
MinTunnelQualityProbeIntervalSec = 1
MaxTunnelQualityProbeIntervalSec = 3600
)
func TunnelQualityProbeIntervalSecondsFromConfigMap(cfg map[string]string) int {
if cfg == nil {
return DefaultTunnelQualityProbeIntervalSec
}
seconds, err := parseTunnelQualityProbeIntervalSeconds(cfg[ConfigTunnelQualityProbeIntervalSec])
if err != nil {
return DefaultTunnelQualityProbeIntervalSec
}
return seconds
}
func NormalizeTunnelQualityProbeIntervalSeconds(value string) (string, error) {
seconds, err := parseTunnelQualityProbeIntervalSeconds(value)
if err != nil {
return "", err
}
return strconv.Itoa(seconds), nil
}
func parseTunnelQualityProbeIntervalSeconds(value string) (int, error) {
trimmed := strings.TrimSpace(value)
if trimmed == "" {
return 0, fmt.Errorf("隧道质量探测间隔不能为空")
}
seconds, err := strconv.Atoi(trimmed)
if err != nil {
return 0, fmt.Errorf("隧道质量探测间隔必须是整数")
}
if seconds < MinTunnelQualityProbeIntervalSec || seconds > MaxTunnelQualityProbeIntervalSec {
return 0, fmt.Errorf(
"隧道质量探测间隔必须在 %d 到 %d 秒之间",
MinTunnelQualityProbeIntervalSec,
MaxTunnelQualityProbeIntervalSec,
)
}
return seconds, nil
}
@@ -0,0 +1,37 @@
package monitoring
import "testing"
func TestTunnelQualityProbeIntervalSecondsFromConfigMap(t *testing.T) {
tests := []struct {
name string
cfg map[string]string
want int
}{
{name: "missing config", cfg: nil, want: DefaultTunnelQualityProbeIntervalSec},
{name: "configured", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "15"}, want: 15},
{name: "invalid", cfg: map[string]string{ConfigTunnelQualityProbeIntervalSec: "0"}, want: DefaultTunnelQualityProbeIntervalSec},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := TunnelQualityProbeIntervalSecondsFromConfigMap(tt.cfg); got != tt.want {
t.Fatalf("interval = %d, want %d", got, tt.want)
}
})
}
}
func TestNormalizeTunnelQualityProbeIntervalSeconds(t *testing.T) {
for _, value := range []string{"1", "15", "3600"} {
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err != nil || got != value {
t.Fatalf("normalize %q = %q, %v", value, got, err)
}
}
for _, value := range []string{"", "0", "3601", "1.5", "abc"} {
if got, err := NormalizeTunnelQualityProbeIntervalSeconds(value); err == nil {
t.Fatalf("normalize %q unexpectedly succeeded with %q", value, got)
}
}
}
@@ -46,7 +46,8 @@ func RenderTable(plan NodePlan) string {
}
targetHost := strings.Trim(strings.TrimSpace(rule.TargetHost), "[]")
for _, protocol := range normalizedProtocols(rule.Protocols) {
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s daddr %s %s dport %d counter comment %q\n",
protocol,
rule.InPort,
family,
targetHost,
@@ -54,7 +55,8 @@ func RenderTable(plan NodePlan) string {
rule.TargetPort,
counterComment(rule.ForwardID, CounterDirectionToTarget, protocol),
))
b.WriteString(fmt.Sprintf(" ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
b.WriteString(fmt.Sprintf(" meta l4proto %s ct original proto-dst %d %s saddr %s %s sport %d counter comment %q\n",
protocol,
rule.InPort,
family,
targetHost,
@@ -62,10 +62,10 @@ func TestRenderTableIncludesForwardAccountingCounters(t *testing.T) {
wantLines := []string{
`tcp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat tcp"`,
`udp dport 12345 counter dnat ip to 198.51.100.20:443 comment "flvx forward:42 dnat udp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
`ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12345 ip saddr 198.51.100.20 tcp sport 443 counter comment "flvx forward:42 from-target tcp"`,
`meta l4proto udp ct original proto-dst 12345 ip daddr 198.51.100.20 udp dport 443 counter comment "flvx forward:42 to-target udp"`,
`meta l4proto udp ct original proto-dst 12345 ip saddr 198.51.100.20 udp sport 443 counter comment "flvx forward:42 from-target udp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
@@ -89,8 +89,8 @@ func TestRenderTableIncludesIPv6ForwardAccountingCounters(t *testing.T) {
got := RenderTable(plan)
wantLines := []string{
`tcp dport 12346 counter dnat ip6 to [2001:db8::20]:8443 comment "flvx forward:43 dnat tcp"`,
`ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
`ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip6 daddr 2001:db8::20 tcp dport 8443 counter comment "flvx forward:43 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip6 saddr 2001:db8::20 tcp sport 8443 counter comment "flvx forward:43 from-target tcp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
@@ -110,8 +110,8 @@ func TestRenderTableAccountingCountersIncludeOriginalPort(t *testing.T) {
got := RenderTable(plan)
wantLines := []string{
`ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12345 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:42 to-target tcp"`,
`meta l4proto tcp ct original proto-dst 12346 ip daddr 198.51.100.20 tcp dport 443 counter comment "flvx forward:43 to-target tcp"`,
}
for _, want := range wantLines {
if !strings.Contains(got, want) {
+43 -8
View File
@@ -26,23 +26,58 @@ func NewSSHRunner() *SSHRunner {
}
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
nft := nftBinary(cfg)
tableName := fmt.Sprintf("flvx_capability_%d", time.Now().UnixNano())
return r.run(ctx, cfg, buildCapabilityCheckCommand(nft, tableName))
}
func buildCapabilityCheckCommand(nft, tableName string) string {
script := RenderTable(NodePlan{
Rules: []Rule{{
ForwardID: 1,
InPort: 12345,
TargetHost: "192.0.2.1",
TargetPort: 443,
Protocols: []string{"tcp", "udp"},
}},
})
script = strings.Replace(script, "table inet flvx {", "table inet "+tableName+" {", 1)
return "set -eu\n" +
"command -v nft >/dev/null 2>&1\n" +
nft + " --version >/dev/null 2>&1\n" +
"tmp=$(mktemp /tmp/flvx-nft-capability-XXXXXX.nft)\n" +
"trap 'rm -f \"$tmp\"' EXIT\n" +
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
"if ! " + nft + " -c -f \"$tmp\"; then\n" +
" echo 'nftables cannot validate the generated FLVX rules' >&2\n" +
" exit 1\n" +
"fi"
}
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
nft := nftBinary(cfg)
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
command := buildApplyCommand(nftBinary(cfg), script)
return r.run(ctx, cfg, command)
}
func buildApplyCommand(nft, script string) string {
return "set -eu\n" +
"tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft)\n" +
"batch=$(mktemp /tmp/flvx-nft-batch-XXXXXX.nft) || { rm -f \"$tmp\"; exit 1; }\n" +
"cleanup() {\n" +
" rm -f \"$tmp\"\n" +
" rm -f \"$tmp\" \"$batch\"\n" +
"}\n" +
"trap cleanup EXIT\n" +
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
nft + " -c -f \"$tmp\"\n" +
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
" " + nft + " delete table inet flvx\n" +
" { printf '%s\\n' 'delete table inet flvx'; cat \"$tmp\"; } > \"$batch\"\n" +
"else\n" +
" cp \"$tmp\" \"$batch\"\n" +
"fi\n" +
nft + " -f \"$tmp\""
return r.run(ctx, cfg, command)
"if ! " + nft + " -c -f \"$batch\"; then\n" +
" echo 'nftables rule validation failed; active rules were preserved' >&2\n" +
" exit 1\n" +
"fi\n" +
nft + " -f \"$batch\""
}
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
@@ -5,10 +5,113 @@ import (
"crypto/rsa"
"crypto/x509"
"encoding/pem"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
)
func TestBuildApplyCommandStopsAfterValidationFailure(t *testing.T) {
dir := t.TempDir()
logPath := filepath.Join(dir, "calls.log")
applyMarker := filepath.Join(dir, "applied")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
echo "$*" >> "` + logPath + `"
if [ "$1" = "list" ]; then
exit 0
fi
if [ "$1" = "-c" ]; then
exit 1
fi
touch "` + applyMarker + `"
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
command := buildApplyCommand(nftPath, "table inet flvx { }")
result := exec.Command("sh", "-c", command)
if err := result.Run(); err == nil {
t.Fatal("expected validation failure")
}
if _, err := os.Stat(applyMarker); !os.IsNotExist(err) {
t.Fatalf("apply ran after validation failure, stat err=%v", err)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("read fake nft calls: %v", err)
}
if strings.Count(string(calls), "-f ") != 1 {
t.Fatalf("expected validation only, got calls:\n%s", calls)
}
}
func TestBuildCapabilityCheckCommandValidatesRenderedRulesWithoutApplying(t *testing.T) {
dir := t.TempDir()
logPath := filepath.Join(dir, "calls.log")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
echo "$*" >> "` + logPath + `"
if [ "$1" = "--version" ] || [ "$1" = "-c" ]; then
exit 0
fi
exit 1
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
command := strings.Replace(buildCapabilityCheckCommand(nftPath, "flvx_capability_test"), "command -v nft", "command -v "+nftPath, 1)
result := exec.Command("sh", "-c", command)
if output, err := result.CombinedOutput(); err != nil {
t.Fatalf("capability command failed: %v: %s", err, output)
}
calls, err := os.ReadFile(logPath)
if err != nil {
t.Fatalf("read fake nft calls: %v", err)
}
if strings.Count(string(calls), "-c -f ") != 1 || strings.Contains(string(calls), "\n-f ") {
t.Fatalf("expected one check-only invocation, got calls:\n%s", calls)
}
if !strings.Contains(command, "table inet flvx_capability_test") || !strings.Contains(command, "meta l4proto tcp ct original proto-dst") {
t.Fatalf("capability check does not contain representative rendered rules:\n%s", command)
}
}
func TestBuildApplyCommandUsesAtomicReplacementBatch(t *testing.T) {
dir := t.TempDir()
batchPath := filepath.Join(dir, "batch.nft")
nftPath := filepath.Join(dir, "nft")
fake := `#!/bin/sh
if [ "$1" = "list" ]; then
exit 0
fi
if [ "$1" = "-f" ]; then
cp "$2" "` + batchPath + `"
fi
exit 0
`
if err := os.WriteFile(nftPath, []byte(fake), 0o755); err != nil {
t.Fatalf("write fake nft: %v", err)
}
script := "table inet flvx {\n chain forward { }\n}"
result := exec.Command("sh", "-c", buildApplyCommand(nftPath, script))
if output, err := result.CombinedOutput(); err != nil {
t.Fatalf("apply command failed: %v: %s", err, output)
}
batch, err := os.ReadFile(batchPath)
if err != nil {
t.Fatalf("read applied batch: %v", err)
}
want := "delete table inet flvx\n" + script + "\n"
if string(batch) != want {
t.Fatalf("atomic batch = %q, want %q", batch, want)
}
}
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
privateKey := mustGeneratePrivateKey(t)
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
@@ -15,14 +15,27 @@ var publicConfigKeys = map[string]struct{}{
"app_favicon": {},
"app_bg_image": {},
"cloudflare_site_key": {},
"is_commercial": {},
"hide_footer_brand": {},
}
var sensitiveConfigKeys = map[string]struct{}{
"jwt_secret": {},
"license_key": {},
"license_expiry": {},
"license_machine_id": {},
"machine_fingerprint": {},
"cloudflare_secret_key": {},
}
var systemManagedConfigKeys = map[string]struct{}{
"license_key": {},
"license_expiry": {},
"license_machine_id": {},
"is_commercial": {},
"machine_fingerprint": {},
}
func PolicyForConfig(name string) ConfigAccessPolicy {
if IsPublicConfigKey(name) {
return ConfigAccessPublic
@@ -43,6 +56,11 @@ func IsSensitiveConfigKey(name string) bool {
return ok
}
func IsSystemManagedConfigKey(name string) bool {
_, ok := systemManagedConfigKeys[normalizeConfigKey(name)]
return ok
}
func FilterSensitiveConfigs(in map[string]string) map[string]string {
if len(in) == 0 {
return map[string]string{}
@@ -57,6 +75,20 @@ func FilterSensitiveConfigs(in map[string]string) map[string]string {
return out
}
func FilterBackupConfigs(in map[string]string) map[string]string {
if len(in) == 0 {
return map[string]string{}
}
out := make(map[string]string, len(in))
for name, value := range in {
if IsSensitiveConfigKey(name) || IsSystemManagedConfigKey(name) {
continue
}
out[name] = value
}
return out
}
func normalizeConfigKey(name string) string {
return strings.ToLower(strings.TrimSpace(name))
}
@@ -13,6 +13,8 @@ func TestConfigPolicy(t *testing.T) {
{name: "app_favicon is public", key: "app_favicon", want: ConfigAccessPublic},
{name: "app_bg_image is public", key: "app_bg_image", want: ConfigAccessPublic},
{name: "cloudflare_site_key is public", key: "cloudflare_site_key", want: ConfigAccessPublic},
{name: "is_commercial is public", key: "is_commercial", want: ConfigAccessPublic},
{name: "hide_footer_brand is public", key: "hide_footer_brand", want: ConfigAccessPublic},
{name: "jwt_secret is sensitive", key: "jwt_secret", want: ConfigAccessSensitive},
{name: "license_key is sensitive", key: "license_key", want: ConfigAccessSensitive},
{name: "cloudflare_secret_key is sensitive", key: "cloudflare_secret_key", want: ConfigAccessSensitive},
@@ -29,14 +31,14 @@ func TestConfigPolicy(t *testing.T) {
}
func TestConfigPolicyHelpers(t *testing.T) {
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "cloudflare_site_key"}
publicKeys := []string{"app_name", "app_logo", "app_favicon", "app_bg_image", "cloudflare_site_key", "is_commercial", "hide_footer_brand"}
for _, key := range publicKeys {
if !IsPublicConfigKey(key) {
t.Fatalf("expected %q to be public", key)
}
}
sensitiveKeys := []string{"jwt_secret", "license_key", "cloudflare_secret_key"}
sensitiveKeys := []string{"jwt_secret", "license_key", "license_expiry", "license_machine_id", "machine_fingerprint", "cloudflare_secret_key"}
for _, key := range sensitiveKeys {
if !IsSensitiveConfigKey(key) {
t.Fatalf("expected %q to be sensitive", key)
@@ -67,3 +69,30 @@ func TestConfigPolicyHelpers(t *testing.T) {
t.Fatal("expected cloudflare_secret_key to be filtered out")
}
}
func TestFilterBackupConfigsOmitsLicenseState(t *testing.T) {
filtered := FilterBackupConfigs(map[string]string{
"app_name": "Brand",
"license_key": "secret-license",
"license_expiry": "never",
"license_machine_id": "machine-id",
"is_commercial": "true",
"machine_fingerprint": "fingerprint",
})
if len(filtered) != 1 || filtered["app_name"] != "Brand" {
t.Fatalf("unexpected backup configs: %+v", filtered)
}
}
func TestSystemManagedConfigKeys(t *testing.T) {
for _, key := range []string{"license_key", "license_expiry", "license_machine_id", "is_commercial", "machine_fingerprint"} {
if !IsSystemManagedConfigKey(key) {
t.Fatalf("expected %s to be system managed", key)
}
}
for _, key := range []string{"jwt_secret", "cloudflare_secret_key", "app_name"} {
if IsSystemManagedConfigKey(key) {
t.Fatalf("did not expect %s to be system managed", key)
}
}
}
+28 -3
View File
@@ -473,6 +473,14 @@ func seedData(db *gorm.DB) {
appNameConfig := model.ViteConfig{ID: 1, Name: "app_name", Value: "flux", Time: 1755147963000}
db.Where("id = ?", 1).FirstOrCreate(&appNameConfig)
now := time.Now().UnixMilli()
for name, value := range map[string]string{
"is_commercial": "false",
"hide_footer_brand": "false",
} {
cfg := model.ViteConfig{Name: name, Value: value, Time: now}
db.Where("name = ?", name).FirstOrCreate(&cfg)
}
}
// ─── User Queries ────────────────────────────────────────────────────
@@ -600,6 +608,23 @@ func (r *Repository) UpsertConfig(name, value string, now int64) error {
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error
}
func (r *Repository) UpsertConfigs(values map[string]string, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
return r.db.Transaction(func(tx *gorm.DB) error {
for name, value := range values {
if err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "name"}},
DoUpdates: clause.AssignmentColumns([]string{"value", "time"}),
}).Create(&model.ViteConfig{Name: name, Value: value, Time: now}).Error; err != nil {
return err
}
}
return nil
})
}
// ─── Announcement Queries ────────────────────────────────────────────
func (r *Repository) GetAnnouncement() (*model.Announcement, error) {
@@ -1967,7 +1992,7 @@ func (r *Repository) ExportAll() (*model.BackupData, error) {
if err != nil {
return nil, fmt.Errorf("export configs failed: %w", err)
}
backup.Configs = FilterSensitiveConfigs(configs)
backup.Configs = FilterBackupConfigs(configs)
return backup, nil
}
@@ -2047,7 +2072,7 @@ func (r *Repository) ExportPartial(types []string) (*model.BackupData, error) {
if err != nil {
return nil, fmt.Errorf("export configs failed: %w", err)
}
backup.Configs = FilterSensitiveConfigs(v)
backup.Configs = FilterBackupConfigs(v)
}
return backup, nil
}
@@ -2839,7 +2864,7 @@ func importPermissions(tx *gorm.DB, permissions []model.PermissionBackup, _ int6
}
func importConfigs(tx *gorm.DB, configs map[string]string, now int64) (int, error) {
configs = FilterSensitiveConfigs(configs)
configs = FilterBackupConfigs(configs)
count := 0
for name, value := range configs {
err := tx.Clauses(clause.OnConflict{
@@ -97,6 +97,10 @@ func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
seedConfig(t, r, "cloudflare_site_key", "site-key")
seedConfig(t, r, "jwt_secret", "jwt-secret")
seedConfig(t, r, "license_key", "license-secret")
seedConfig(t, r, "license_expiry", "2030-01-02T00:00:00.000Z")
seedConfig(t, r, "license_machine_id", "machine-id")
seedConfig(t, r, "is_commercial", "true")
seedConfig(t, r, "machine_fingerprint", "machine-fingerprint")
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-secret")
for _, tc := range []struct {
@@ -117,7 +121,7 @@ func TestExportAllOmitsSensitiveConfigs(t *testing.T) {
if backup.Configs["cloudflare_site_key"] != "site-key" {
t.Fatalf("expected public config in export, got %+v", backup.Configs)
}
for _, key := range []string{"jwt_secret", "license_key", "cloudflare_secret_key"} {
for _, key := range []string{"jwt_secret", "license_key", "license_expiry", "license_machine_id", "is_commercial", "machine_fingerprint", "cloudflare_secret_key"} {
if _, ok := backup.Configs[key]; ok {
t.Fatalf("expected %s to be omitted from export, got %+v", key, backup.Configs)
}
@@ -136,12 +140,20 @@ func TestImportIgnoresSensitiveConfigs(t *testing.T) {
seedConfig(t, r, "app_name", "before")
seedConfig(t, r, "jwt_secret", "jwt-before")
seedConfig(t, r, "license_key", "license-before")
seedConfig(t, r, "license_expiry", "expiry-before")
seedConfig(t, r, "license_machine_id", "machine-before")
seedConfig(t, r, "is_commercial", "true")
seedConfig(t, r, "machine_fingerprint", "fingerprint-before")
seedConfig(t, r, "cloudflare_secret_key", "cloudflare-before")
backup := &model.BackupData{Configs: map[string]string{
"app_name": "after",
"jwt_secret": "jwt-after",
"license_key": "license-after",
"license_expiry": "expiry-after",
"license_machine_id": "machine-after",
"is_commercial": "false",
"machine_fingerprint": "fingerprint-after",
"cloudflare_secret_key": "cloudflare-after",
}}
@@ -156,6 +168,10 @@ func TestImportIgnoresSensitiveConfigs(t *testing.T) {
assertConfigValue(t, r, "app_name", "after")
assertConfigValue(t, r, "jwt_secret", "jwt-before")
assertConfigValue(t, r, "license_key", "license-before")
assertConfigValue(t, r, "license_expiry", "expiry-before")
assertConfigValue(t, r, "license_machine_id", "machine-before")
assertConfigValue(t, r, "is_commercial", "true")
assertConfigValue(t, r, "machine_fingerprint", "fingerprint-before")
assertConfigValue(t, r, "cloudflare_secret_key", "cloudflare-before")
}
@@ -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)
}
}
@@ -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 ""
@@ -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")
}
}
+82 -3
View File
@@ -3,6 +3,7 @@ package config
import (
"encoding/json"
"io"
"reflect"
"sync"
"time"
@@ -30,9 +31,87 @@ func Global() *Config {
globalMux.RLock()
defer globalMux.RUnlock()
cfg := &Config{}
*cfg = *global
return cfg
return cloneConfig(global)
}
// cloneConfig returns a detached snapshot of the runtime config. Runtime
// commands mutate slices, pointers and metadata maps, so a shallow struct copy
// is not sufficient once readers and persistence run concurrently.
func cloneConfig(c *Config) *Config {
if c == nil {
return nil
}
v := cloneConfigValue(reflect.ValueOf(c))
if !v.IsValid() || v.IsNil() {
return nil
}
return v.Interface().(*Config)
}
func cloneConfigValue(v reflect.Value) reflect.Value {
if !v.IsValid() {
return reflect.Value{}
}
switch v.Kind() {
case reflect.Interface:
if v.IsNil() {
return reflect.Zero(v.Type())
}
cloned := cloneConfigValue(v.Elem())
out := reflect.New(v.Type()).Elem()
if cloned.IsValid() && cloned.Type().AssignableTo(v.Type()) {
out.Set(cloned)
} else if cloned.IsValid() && cloned.Type().Implements(v.Type()) {
out.Set(cloned)
} else if cloned.IsValid() {
out.Set(cloned)
}
return out
case reflect.Pointer:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.New(v.Type().Elem())
out.Elem().Set(cloneConfigValue(v.Elem()))
return out
case reflect.Map:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeMapWithSize(v.Type(), v.Len())
iter := v.MapRange()
for iter.Next() {
out.SetMapIndex(cloneConfigValue(iter.Key()), cloneConfigValue(iter.Value()))
}
return out
case reflect.Slice:
if v.IsNil() {
return reflect.Zero(v.Type())
}
out := reflect.MakeSlice(v.Type(), v.Len(), v.Len())
for i := 0; i < v.Len(); i++ {
out.Index(i).Set(cloneConfigValue(v.Index(i)))
}
return out
case reflect.Array:
out := reflect.New(v.Type()).Elem()
for i := 0; i < v.Len(); i++ {
out.Index(i).Set(cloneConfigValue(v.Index(i)))
}
return out
case reflect.Struct:
out := reflect.New(v.Type()).Elem()
out.Set(v)
for i := 0; i < v.NumField(); i++ {
if out.Field(i).CanSet() && v.Field(i).CanInterface() {
out.Field(i).Set(cloneConfigValue(v.Field(i)))
}
}
return out
default:
return v
}
}
func Set(c *Config) {
+121
View File
@@ -0,0 +1,121 @@
package config
import (
"encoding/json"
"path/filepath"
"sync"
"testing"
)
func TestGlobalReturnsDetachedSnapshot(t *testing.T) {
original := Global()
t.Cleanup(func() { Set(original) })
Set(&Config{Services: []*ServiceConfig{{
Name: "snapshot-service",
Metadata: map[string]any{"paused": false},
Handler: &HandlerConfig{
Type: "relay",
Metadata: map[string]any{"retries": 2},
},
}}})
snapshot := Global()
snapshot.Services[0].Name = "changed"
snapshot.Services[0].Metadata["paused"] = true
snapshot.Services[0].Handler.Metadata["retries"] = 9
current := Global()
if current.Services[0].Name != "snapshot-service" {
t.Fatalf("snapshot mutated global service name: %q", current.Services[0].Name)
}
if paused, _ := current.Services[0].Metadata["paused"].(bool); paused {
t.Fatalf("snapshot mutated global service metadata")
}
if retries, _ := current.Services[0].Handler.Metadata["retries"].(int); retries != 2 {
t.Fatalf("snapshot mutated nested handler metadata: %v", current.Services[0].Handler.Metadata["retries"])
}
}
func TestConcurrentGlobalSnapshotAndUpdate(t *testing.T) {
original := Global()
t.Cleanup(func() { Set(original) })
Set(&Config{Services: []*ServiceConfig{{
Name: "concurrent-service",
Metadata: map[string]any{"generation": 0},
}}})
var wg sync.WaitGroup
for worker := 0; worker < 4; worker++ {
wg.Add(1)
go func(worker int) {
defer wg.Done()
for i := 0; i < 500; i++ {
if err := OnUpdate(func(c *Config) error {
c.Services[0] = &ServiceConfig{
Name: "concurrent-service",
Metadata: map[string]any{"generation": worker*500 + i},
}
return nil
}); err != nil {
t.Errorf("OnUpdate: %v", err)
return
}
}
}(worker)
}
for i := 0; i < 2000; i++ {
if _, err := json.Marshal(Global()); err != nil {
t.Fatalf("marshal snapshot: %v", err)
}
}
wg.Wait()
}
func TestConcurrentPersistProducesValidConfig(t *testing.T) {
original := Global()
originalPath := PersistPath()
persistMu.Lock()
originalEnabled := persistEnable
persistMu.Unlock()
t.Cleanup(func() {
Set(original)
SetPersistPath(originalPath)
persistMu.Lock()
persistEnable = originalEnabled
persistMu.Unlock()
})
path := filepath.Join(t.TempDir(), "gost.json")
SetPersistPath(path)
EnablePersist()
Set(&Config{Services: []*ServiceConfig{{Name: "persist-service"}}})
var wg sync.WaitGroup
for worker := 0; worker < 4; worker++ {
wg.Add(1)
go func(worker int) {
defer wg.Done()
for i := 0; i < 50; i++ {
if err := OnUpdate(func(c *Config) error {
c.Services[0].Metadata = map[string]any{"generation": worker*50 + i}
return nil
}); err != nil {
t.Errorf("persist update: %v", err)
return
}
}
}(worker)
}
wg.Wait()
var persisted Config
if err := persisted.ReadFile(path); err != nil {
t.Fatalf("read persisted config: %v", err)
}
if len(persisted.Services) != 1 || persisted.Services[0] == nil || persisted.Services[0].Name != "persist-service" {
t.Fatalf("unexpected persisted services: %#v", persisted.Services)
}
}
+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)
}
}
+1 -1
View File
@@ -42,9 +42,9 @@ func EnablePersist() {
// persist writes the current global config to the configured file atomically.
func persist() error {
persistMu.Lock()
defer persistMu.Unlock()
path := persistPath
enabled := persistEnable
persistMu.Unlock()
if !enabled || path == "" {
return nil
+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")
}
}
+117
View File
@@ -0,0 +1,117 @@
package socket
import (
"net"
"sync"
"testing"
"time"
coreservice "github.com/go-gost/core/service"
"github.com/go-gost/x/config"
"github.com/go-gost/x/registry"
)
type blockingCommandService struct {
started chan struct{}
release chan struct{}
startedOnce sync.Once
}
func (s *blockingCommandService) Serve() error { return nil }
func (s *blockingCommandService) Addr() net.Addr { return nil }
func (s *blockingCommandService) Close() error {
s.startedOnce.Do(func() { close(s.started) })
<-s.release
return nil
}
func TestMutationCommandsAreSerialized(t *testing.T) {
mutations := []string{
"AddService", "UpdateService", "DeleteService", "PauseService", "ResumeService",
"AddChains", "UpdateChains", "DeleteChains",
"AddLimiters", "UpdateLimiters", "DeleteLimiters",
"AddCLimiters", "UpdateCLimiters", "DeleteCLimiters",
"SetProtocol", "UpgradeAgent", "RollbackAgent", "reload",
}
for _, command := range mutations {
if !isMutationCommand(command) {
t.Fatalf("expected %s to use the serialized mutation queue", command)
}
}
}
func TestReadOnlyCommandsRemainBoundedAsync(t *testing.T) {
for _, command := range []string{"TcpPing", "UdpPing", "ServiceMonitorCheck"} {
if isMutationCommand(command) {
t.Fatalf("expected %s to remain a read-only command", command)
}
}
}
func TestCommandResponseTypePreservesRequestContract(t *testing.T) {
if got := commandResponseType("UpdateService"); got != "UpdateServiceResponse" {
t.Fatalf("unexpected response type: %s", got)
}
if got := commandResponseType(""); got != "UnknownCommandResponse" {
t.Fatalf("unexpected empty command response type: %s", got)
}
}
func TestMutationQueueExecutesCommandsInArrivalOrder(t *testing.T) {
originalConfig := config.Global()
t.Cleanup(func() { config.Set(originalConfig) })
firstName := "mutation_queue_first_tdd"
secondName := "mutation_queue_second_tdd"
first := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
second := &blockingCommandService{started: make(chan struct{}), release: make(chan struct{})}
for _, name := range []string{firstName, secondName} {
registry.ServiceRegistry().Unregister(name)
}
t.Cleanup(func() {
select {
case <-first.release:
default:
close(first.release)
}
select {
case <-second.release:
default:
close(second.release)
}
registry.ServiceRegistry().Unregister(firstName)
registry.ServiceRegistry().Unregister(secondName)
})
if err := registry.ServiceRegistry().Register(firstName, coreservice.Service(first)); err != nil {
t.Fatalf("register first service: %v", err)
}
if err := registry.ServiceRegistry().Register(secondName, coreservice.Service(second)); err != nil {
t.Fatalf("register second service: %v", err)
}
config.Set(&config.Config{Services: []*config.ServiceConfig{{Name: firstName}, {Name: secondName}}})
reporter := NewWebSocketReporter("", "mutation-queue-test-secret")
go reporter.runMutationCommands()
t.Cleanup(reporter.Stop)
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{firstName}}})
reporter.dispatchCommand(CommandMessage{Type: "DeleteService", Data: map[string]any{"services": []string{secondName}}})
select {
case <-first.started:
case <-time.After(time.Second):
t.Fatal("first mutation did not start")
}
select {
case <-second.started:
t.Fatal("second mutation started before first mutation completed")
case <-time.After(100 * time.Millisecond):
}
close(first.release)
select {
case <-second.started:
case <-time.After(time.Second):
t.Fatal("second mutation did not start after first mutation completed")
}
close(second.release)
}
+82 -18
View File
@@ -15,6 +15,11 @@ import (
xservice "github.com/go-gost/x/service"
)
type serviceReplacement struct {
config config.ServiceConfig
oldConfig *config.ServiceConfig
}
func createServices(req createServicesRequest) error {
if len(req.Data) == 0 {
@@ -95,11 +100,11 @@ func updateServices(req updateServicesRequest) error {
req.Data[i].Name = name
}
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)
changedServices := make([]struct {
config config.ServiceConfig
service service.Service
}, 0, len(req.Data))
// 第二阶段:逐个更新服务(Upsert模式:存在则更新,不存在则创建)。
// 配置变更命令由 WebSocket reporter 串行调度,但这里仍保留完整回滚,
// 避免新配置解析或监听失败后旧服务永久消失。
originalConfig := config.Global()
changedServices := make([]serviceReplacement, 0, len(req.Data))
for i := range req.Data {
serviceConfig := &req.Data[i]
name := serviceConfig.Name
@@ -107,33 +112,48 @@ func updateServices(req updateServicesRequest) error {
continue
}
// 1. 获取旧服务
old := registry.ServiceRegistry().Get(name)
var oldConfig *config.ServiceConfig
if originalConfig != nil {
for _, current := range originalConfig.Services {
if current != nil && strings.TrimSpace(current.Name) == name {
oldConfig = current
break
}
}
}
// 2. 关闭旧服务 (如果存在)
if old != nil {
// 3. 从注册表移除旧服务;registry 会负责关闭旧服务。
// 1. 关闭并移除旧服务(如果存在)。同名监听必须先释放端口,
// 才能创建新 listener。
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
// 4. 解析新服务配置
// 2. 解析新服务配置。
svc, err := parser.ParseService(serviceConfig)
if err != nil {
rollbackErr := restoreServiceRuntime(name, oldConfig)
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
if rollbackErr != nil {
return fmt.Errorf("create service %s failed: %v; restore previous service failed: %w", name, err, rollbackErr)
}
return errors.New("create service " + name + " failed: " + err.Error())
}
changedServices = append(changedServices, struct {
config config.ServiceConfig
service service.Service
}{*serviceConfig, svc})
// 5. 注册新服务
// 3. 注册并启动新服务。
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
svc.Close()
rollbackErr := restoreServiceRuntime(name, oldConfig)
rollbackErr = errors.Join(rollbackErr, rollbackServiceReplacements(changedServices))
if rollbackErr != nil {
return fmt.Errorf("service %s already exists; restore previous service failed: %w", name, rollbackErr)
}
return errors.New("service " + name + " already exists")
}
// 6. 启动新服务
go svc.Serve()
changedServices = append(changedServices, serviceReplacement{
config: *serviceConfig,
oldConfig: oldConfig,
})
}
if len(changedServices) == 0 {
return nil
@@ -158,12 +178,56 @@ func updateServices(req updateServicesRequest) error {
}
return nil
}); err != nil {
config.Set(originalConfig)
if rollbackErr := rollbackServiceReplacements(changedServices); rollbackErr != nil {
return fmt.Errorf("%w; restore previous services failed: %v", err, rollbackErr)
}
return err
}
return nil
}
func restoreServiceRuntime(name string, serviceConfig *config.ServiceConfig) error {
name = strings.TrimSpace(name)
if name == "" || serviceConfig == nil {
return nil
}
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
cfgCopy := *serviceConfig
cfgCopy.Name = name
svc, err := parser.ParseService(&cfgCopy)
if err != nil {
return err
}
if err := registry.ServiceRegistry().Register(name, svc); err != nil {
svc.Close()
return err
}
go svc.Serve()
return nil
}
func rollbackServiceReplacements(replacements []serviceReplacement) error {
var rollbackErr error
for i := len(replacements) - 1; i >= 0; i-- {
name := strings.TrimSpace(replacements[i].config.Name)
if name == "" {
continue
}
if registry.ServiceRegistry().Get(name) != nil {
registry.ServiceRegistry().Unregister(name)
}
if err := restoreServiceRuntime(name, replacements[i].oldConfig); err != nil {
rollbackErr = errors.Join(rollbackErr, fmt.Errorf("restore service %s: %w", name, err))
}
}
return rollbackErr
}
func serviceConfigUnchanged(name string, next config.ServiceConfig) bool {
cfg := config.Global()
if cfg == nil {
+39
View File
@@ -7,6 +7,8 @@ import (
corelogger "github.com/go-gost/core/logger"
"github.com/go-gost/core/service"
"github.com/go-gost/x/config"
_ "github.com/go-gost/x/handler/auto"
_ "github.com/go-gost/x/listener/tcp"
xlogger "github.com/go-gost/x/logger"
"github.com/go-gost/x/registry"
)
@@ -15,6 +17,43 @@ type recordingService struct {
closed int
}
func TestUpdateServicesParseFailureRestoresPreviousRuntime(t *testing.T) {
corelogger.SetDefault(xlogger.Nop())
name := "restore_after_failed_update_tdd"
existing := &recordingService{}
registry.ServiceRegistry().Unregister(name)
t.Cleanup(func() { registry.ServiceRegistry().Unregister(name) })
if err := registry.ServiceRegistry().Register(name, service.Service(existing)); err != nil {
t.Fatalf("register existing service: %v", err)
}
originalConfig := config.Global()
t.Cleanup(func() { config.Set(originalConfig) })
serviceConfig := config.ServiceConfig{Name: name, Addr: "127.0.0.1:0"}
config.Set(&config.Config{Services: []*config.ServiceConfig{&serviceConfig}})
invalid := serviceConfig
invalid.Listener = &config.ListenerConfig{Type: "listener-does-not-exist"}
if err := updateServices(updateServicesRequest{Data: []config.ServiceConfig{invalid}}); err == nil {
t.Fatalf("expected invalid service update to fail")
}
if existing.closed != 1 {
t.Fatalf("expected old runtime to be closed once, got %d", existing.closed)
}
if registry.ServiceRegistry().Get(name) == nil {
t.Fatalf("expected previous runtime to be restored")
}
cfg := config.Global()
if len(cfg.Services) != 1 || cfg.Services[0] == nil || cfg.Services[0].Name != name {
t.Fatalf("expected previous config to remain, got %#v", cfg.Services)
}
if cfg.Services[0].Listener != nil {
t.Fatalf("expected previous listener config to remain unchanged")
}
}
func (s *recordingService) Serve() error { return nil }
func (s *recordingService) Addr() net.Addr { return nil }
func (s *recordingService) Close() error {
+121 -5
View File
@@ -16,6 +16,7 @@ import (
"os"
"os/exec"
"runtime"
"runtime/debug"
"strconv"
"strings"
"sync"
@@ -151,6 +152,9 @@ const (
initialBackoff = 2 * time.Second // 重连初始退避
maxBackoff = 2 * time.Minute // 重连最大退避
defaultMetricReportInterval = 5 * time.Second
maxConcurrentTCPPings = 8
maxConcurrentReadCommands = 16
maxQueuedMutationCommands = 256
)
type WebSocketReporter struct {
@@ -172,6 +176,9 @@ type WebSocketReporter struct {
connecting bool // 正在连接状态
connMutex sync.Mutex // 连接状态锁
aesCrypto *crypto.AESCrypto // AES加密器
tcpPingSem chan struct{} // 限制诊断探测并发,避免离线目标耗尽连接
readCommandSem chan struct{} // 限制只读命令并发,避免诊断请求耗尽资源
mutationQueue chan CommandMessage
}
var wsDial = func(dialer *websocket.Dialer, rawURL string) (*websocket.Conn, *http.Response, error) {
@@ -201,11 +208,37 @@ func NewWebSocketReporter(serverURL string, secret string) *WebSocketReporter {
connected: false,
connecting: false,
aesCrypto: aesCrypto,
tcpPingSem: make(chan struct{}, maxConcurrentTCPPings),
readCommandSem: make(chan struct{}, maxConcurrentReadCommands),
mutationQueue: make(chan CommandMessage, maxQueuedMutationCommands),
}
}
func (w *WebSocketReporter) tryAcquireTCPPingSlot() bool {
if w == nil || w.tcpPingSem == nil {
return false
}
select {
case w.tcpPingSem <- struct{}{}:
return true
default:
return false
}
}
func (w *WebSocketReporter) releaseTCPPingSlot() {
if w == nil || w.tcpPingSem == nil {
return
}
select {
case <-w.tcpPingSem:
default:
}
}
// Start 启动WebSocket报告器
func (w *WebSocketReporter) Start() {
go w.runMutationCommands()
go w.run()
}
@@ -752,8 +785,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
}
if cmdMsg.Type != "call" {
// 所有命令统一异步执行,避免阻塞消息接收循环
go w.routeCommand(cmdMsg)
w.dispatchCommand(cmdMsg)
}
} else {
// 处理普通消息
@@ -764,8 +796,7 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
return
}
if cmdMsg.Type != "call" {
// 所有命令统一异步执行,避免阻塞消息接收循环
go w.routeCommand(cmdMsg)
w.dispatchCommand(cmdMsg)
}
}
@@ -774,6 +805,86 @@ func (w *WebSocketReporter) handleReceivedMessage(messageType int, message []byt
}
}
// dispatchCommand keeps all runtime mutations ordered while allowing bounded
// concurrency for read-only diagnostics. Mutations share process-wide
// registries and configuration, so running them concurrently can corrupt the
// persisted config or interleave service lifecycle operations.
func (w *WebSocketReporter) dispatchCommand(cmd CommandMessage) {
if isMutationCommand(cmd.Type) {
select {
case w.mutationQueue <- cmd:
case <-w.ctx.Done():
w.sendCommandFailure(cmd, "Agent is shutting down")
default:
w.sendCommandFailure(cmd, "运行时配置命令队列已满,请稍后重试")
}
return
}
select {
case w.readCommandSem <- struct{}{}:
go func() {
defer func() { <-w.readCommandSem }()
w.routeCommandSafely(cmd)
}()
case <-w.ctx.Done():
w.sendCommandFailure(cmd, "Agent is shutting down")
default:
w.sendCommandFailure(cmd, "只读命令并发过多,请稍后重试")
}
}
func (w *WebSocketReporter) runMutationCommands() {
for {
select {
case <-w.ctx.Done():
return
case cmd := <-w.mutationQueue:
w.routeCommandSafely(cmd)
}
}
}
func (w *WebSocketReporter) routeCommandSafely(cmd CommandMessage) {
defer func() {
if recovered := recover(); recovered != nil {
fmt.Printf("❌ 命令处理 panic: type=%s panic=%v\n%s", cmd.Type, recovered, debug.Stack())
w.sendCommandFailure(cmd, fmt.Sprintf("命令处理异常: %v", recovered))
}
}()
w.routeCommand(cmd)
}
func (w *WebSocketReporter) sendCommandFailure(cmd CommandMessage, message string) {
w.sendResponse(CommandResponse{
Type: commandResponseType(cmd.Type),
Success: false,
Message: message,
RequestId: cmd.RequestId,
})
}
func commandResponseType(commandType string) string {
commandType = strings.TrimSpace(commandType)
if commandType == "" {
return "UnknownCommandResponse"
}
return commandType + "Response"
}
func isMutationCommand(commandType string) bool {
switch strings.ToLower(strings.TrimSpace(commandType)) {
case "addservice", "updateservice", "deleteservice", "pauseservice", "resumeservice",
"addchains", "updatechains", "deletechains",
"addlimiters", "updatelimiters", "deletelimiters",
"addclimiters", "updateclimiters", "deleteclimiters",
"setprotocol", "upgradeagent", "rollbackagent", "reload":
return true
default:
return false
}
}
// routeCommand 路由命令到对应的处理函数
func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
jsonBytes, errs := json.Marshal(cmd)
@@ -840,9 +951,14 @@ func (w *WebSocketReporter) routeCommand(cmd CommandMessage) {
// TCP Ping 诊断命令(只读,不需要保存配置)
case "TcpPing":
response.Type = "TcpPingResponse"
if !w.tryAcquireTCPPingSlot() {
err = fmt.Errorf("TCP探测任务过多,请稍后重试")
break
}
defer w.releaseTCPPingSlot()
var tcpPingResult TcpPingResponse
tcpPingResult, err = w.handleTcpPing(cmd.Data)
response.Type = "TcpPingResponse"
response.Data = tcpPingResult
// needSaveConfig = false (默认值)
@@ -148,6 +148,25 @@ func TestNewWebSocketReporterUsesReducedMetricInterval(t *testing.T) {
}
}
func TestWebSocketReporterLimitsConcurrentTCPPings(t *testing.T) {
reporter := &WebSocketReporter{tcpPingSem: make(chan struct{}, maxConcurrentTCPPings)}
for i := 0; i < maxConcurrentTCPPings; i++ {
if !reporter.tryAcquireTCPPingSlot() {
t.Fatalf("expected TCP ping slot %d to be available", i)
}
}
if reporter.tryAcquireTCPPingSlot() {
t.Fatalf("expected TCP ping concurrency limit at %d", maxConcurrentTCPPings)
}
for i := 0; i < maxConcurrentTCPPings; i++ {
reporter.releaseTCPPingSlot()
}
if !reporter.tryAcquireTCPPingSlot() {
t.Fatalf("expected released TCP ping slot to be reusable")
}
reporter.releaseTCPPingSlot()
}
func TestFormatWebSocketDialErrorIncludesHTTPStatus(t *testing.T) {
err := errors.New("websocket: bad handshake")
resp := &http.Response{
+290 -42
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,205 @@ write_flux_agent_config() {
"$(json_escape "$SECRET")" > "$path"
}
ensure_service_manager() {
if [[ -n "$SERVICE_MANAGER" ]]; then
case "$SERVICE_MANAGER" in
systemd|openrc)
return 0
;;
*)
echo "❌ 不支持的服务管理器: $SERVICE_MANAGER" >&2
return 1
;;
esac
fi
# Alpine uses OpenRC even if a systemctl compatibility command happens to be installed.
if [[ -f /etc/alpine-release ]]; then
if command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
SERVICE_MANAGER="openrc"
return 0
fi
elif command -v systemctl >/dev/null 2>&1 && [[ -d /run/systemd/system ]]; then
SERVICE_MANAGER="systemd"
return 0
elif command -v rc-service >/dev/null 2>&1 && command -v rc-update >/dev/null 2>&1; then
SERVICE_MANAGER="openrc"
return 0
fi
echo "❌ 未检测到受支持的服务管理器(systemd 或 OpenRC)。" >&2
return 1
}
flux_agent_service_exists() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
[[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" ]] || \
systemctl list-units --full -all 2>/dev/null | grep -Fq "flux_agent.service"
;;
openrc)
[[ -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]
;;
esac
}
stop_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl stop flux_agent 2>/dev/null || true
;;
openrc)
rc-service flux_agent stop 2>/dev/null || true
;;
esac
}
disable_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl disable flux_agent 2>/dev/null || true
;;
openrc)
rc-update del flux_agent default 2>/dev/null || true
;;
esac
}
write_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
mkdir -p "$(dirname "$FLUX_AGENT_SYSTEMD_SERVICE_FILE")"
cat > "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" <<EOF
[Unit]
Description=Flux_agent Proxy Service
After=network.target
[Service]
WorkingDirectory=$INSTALL_DIR
ExecStart=$INSTALL_DIR/flux_agent
Environment=GODEBUG=disablethp=1
Restart=on-failure
StandardOutput=null
StandardError=null
[Install]
WantedBy=multi-user.target
EOF
;;
openrc)
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
cat > "$FLUX_AGENT_OPENRC_SERVICE_FILE" <<EOF
#!/sbin/openrc-run
name="flux_agent"
description="Flux_agent Proxy Service"
command="$INSTALL_DIR/flux_agent"
directory="$INSTALL_DIR"
command_background="yes"
pidfile="/run/\${RC_SVCNAME}.pid"
output_log="/dev/null"
error_log="/dev/null"
depend() {
need net
}
EOF
chmod +x "$FLUX_AGENT_OPENRC_SERVICE_FILE"
;;
esac
}
enable_and_start_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl daemon-reload
systemctl enable flux_agent
systemctl start flux_agent
;;
openrc)
rc-update add flux_agent default
rc-service flux_agent start
;;
esac
}
start_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl start flux_agent
;;
openrc)
rc-service flux_agent start
;;
esac
}
flux_agent_service_is_active() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl is-active --quiet flux_agent
;;
openrc)
rc-service flux_agent status >/dev/null 2>&1
;;
esac
}
flux_agent_service_status() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
systemctl is-active flux_agent 2>/dev/null || true
;;
openrc)
rc-service flux_agent status 2>/dev/null || true
;;
esac
}
remove_flux_agent_service() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
rm -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE"
systemctl daemon-reload 2>/dev/null || true
;;
openrc)
rm -f "$FLUX_AGENT_OPENRC_SERVICE_FILE"
;;
esac
}
flux_agent_service_status_hint() {
ensure_service_manager || return 1
case "$SERVICE_MANAGER" in
systemd)
echo "systemctl status flux_agent --no-pager"
;;
openrc)
echo "rc-service flux_agent status"
;;
esac
}
cleanup_legacy_gost_installation() {
local matched_service_files=()
local service_file=""
@@ -278,7 +541,8 @@ cleanup_legacy_gost_installation() {
return 0
fi
if systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
if [[ "$SERVICE_MANAGER" == "systemd" ]] && \
systemctl list-units --full -all 2>/dev/null | grep -Fq "gost.service"; then
systemctl stop gost 2>/dev/null || true
systemctl disable gost 2>/dev/null || true
fi
@@ -297,7 +561,7 @@ cleanup_legacy_gost_installation() {
rm -f "$LEGACY_GOST_CONFIG_DIR/gost"
fi
if [[ "$removed_service_file" == "1" ]]; then
if [[ "$removed_service_file" == "1" && "$SERVICE_MANAGER" == "systemd" ]]; then
systemctl daemon-reload 2>/dev/null || true
fi
}
@@ -341,19 +605,21 @@ install_flux_agent() {
get_config_params
# 检查并安装 tcpkill
# 检查并安装 tcpkill
check_and_install_tcpkill
ensure_service_manager || exit 1
mkdir -p "$INSTALL_DIR"
local tmp_binary="$INSTALL_DIR/flux_agent.new"
# 停止并禁用已有服务
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
if flux_agent_service_exists; then
echo "🔍 检测到已存在的flux_agent服务"
systemctl stop flux_agent 2>/dev/null && echo "🛑 停止服务"
systemctl disable flux_agent 2>/dev/null && echo "🚫 禁用自启"
stop_flux_agent_service
echo "🛑 停止服务"
disable_flux_agent_service
echo "🚫 禁用自启"
fi
# 下载 flux_agent
@@ -392,38 +658,21 @@ EOF
# 加强权限
chmod 600 "$INSTALL_DIR"/*.json
# 创建 systemd 服务
SERVICE_FILE="/etc/systemd/system/flux_agent.service"
cat > "$SERVICE_FILE" <<EOF
[Unit]
Description=Flux_agent Proxy Service
After=network.target
[Service]
WorkingDirectory=$INSTALL_DIR
ExecStart=$INSTALL_DIR/flux_agent
Restart=on-failure
StandardOutput=null
StandardError=null
[Install]
WantedBy=multi-user.target
EOF
# 创建 systemd 或 OpenRC 服务
write_flux_agent_service
# 启动服务
systemctl daemon-reload
systemctl enable flux_agent
systemctl start flux_agent
enable_and_start_flux_agent_service
# 检查状态
echo "🔄 检查服务状态..."
if systemctl is-active --quiet flux_agent; then
if flux_agent_service_is_active; then
echo "✅ 安装完成,flux_agent服务已启动并设置为开机启动。"
echo "📁 配置目录: $INSTALL_DIR"
echo "🔧 服务状态: $(systemctl is-active flux_agent)"
echo "🔧 服务状态: $(flux_agent_service_status)"
else
echo "❌ flux_agent服务启动失败,请执行以下命令查看状态:"
echo "systemctl status flux_agent --no-pager"
flux_agent_service_status_hint
fi
}
@@ -443,6 +692,7 @@ update_flux_agent() {
# 检查并安装 tcpkill
check_and_install_tcpkill
ensure_service_manager || return 1
# 先下载新版本
echo "⬇️ 下载最新版本..."
@@ -455,9 +705,9 @@ update_flux_agent() {
cleanup_legacy_gost_installation
# 停止服务
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
if flux_agent_service_exists; then
echo "🛑 停止 flux_agent 服务..."
systemctl stop flux_agent
stop_flux_agent_service
fi
# 替换文件
@@ -469,7 +719,7 @@ update_flux_agent() {
# 重启服务
echo "🔄 重启服务..."
systemctl start flux_agent
start_flux_agent_service
echo "✅ 更新完成,服务已重新启动。"
}
@@ -477,6 +727,7 @@ update_flux_agent() {
# 卸载功能
uninstall_flux_agent() {
echo "🗑️ 开始卸载 flux_agent..."
ensure_service_manager || return 1
read -p "确认卸载 flux_agent 吗?此操作将删除所有相关文件 (y/N): " confirm
if [[ "$confirm" != "y" && "$confirm" != "Y" ]]; then
@@ -485,15 +736,15 @@ uninstall_flux_agent() {
fi
# 停止并禁用服务
if systemctl list-units --full -all | grep -Fq "flux_agent.service"; then
if flux_agent_service_exists; then
echo "🛑 停止并禁用服务..."
systemctl stop flux_agent 2>/dev/null
systemctl disable flux_agent 2>/dev/null
stop_flux_agent_service
disable_flux_agent_service
fi
# 删除服务文件
if [[ -f "/etc/systemd/system/flux_agent.service" ]]; then
rm -f "/etc/systemd/system/flux_agent.service"
if [[ -f "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || -f "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]]; then
remove_flux_agent_service
echo "🧹 删除服务文件"
fi
@@ -503,9 +754,6 @@ uninstall_flux_agent() {
echo "🧹 删除安装目录: $INSTALL_DIR"
fi
# 重载 systemd
systemctl daemon-reload
echo "✅ 卸载完成"
}
+319 -48
View File
@@ -16,6 +16,9 @@ PINNED_VERSION=""
# 镜像加速配置(可由面板传入或交互式询问)
PROXY_ENABLED="${PROXY_ENABLED:-}"
PROXY_URL="${PROXY_URL:-}"
DEFAULT_PANEL_BACKEND_CONTAINER="flux-panel-backend"
DEFAULT_PANEL_POSTGRES_CONTAINER="flux-panel-postgres"
PANEL_SCRIPT_PATH="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)/$(basename "${BASH_SOURCE[0]}")"
# 镜像加速
maybe_proxy_url() {
@@ -304,8 +307,8 @@ get_env_var() {
get_current_db_type() {
local db_type database_url
db_type=$(get_env_var "DB_TYPE")
database_url=$(get_env_var "DATABASE_URL")
db_type=$(get_env_var "DB_TYPE" || true)
database_url=$(get_env_var "DATABASE_URL" || true)
if [[ "$db_type" == "sqlite" ]]; then
echo "sqlite"
@@ -316,13 +319,289 @@ get_current_db_type() {
fi
}
get_container_compose_label() {
local container="$1"
local label="$2"
local value
value=$(docker inspect -f "{{ index .Config.Labels \"$label\" }}" "$container" 2>/dev/null || true)
if [[ "$value" == "<no value>" ]]; then
value=""
fi
printf '%s' "$value"
}
resolve_panel_deployment() {
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
local requested_dir="${PANEL_DEPLOY_DIR:-}"
local label_dir label_project deploy_dir project_name
label_dir=$(get_container_compose_label "$backend_container" "com.docker.compose.project.working_dir")
label_project=$(get_container_compose_label "$backend_container" "com.docker.compose.project")
if [[ -n "$requested_dir" ]]; then
deploy_dir="$requested_dir"
elif [[ -n "$label_dir" ]]; then
deploy_dir="$label_dir"
elif [[ -f ".env" && -f "docker-compose.yml" ]]; then
deploy_dir=$(pwd -P)
else
echo "❌ 无法识别面板部署目录:未找到容器 Compose 标签,当前目录也没有完整部署配置"
return 1
fi
if [[ ! -d "$deploy_dir" ]]; then
echo "❌ 面板部署目录不存在:$deploy_dir"
return 1
fi
deploy_dir=$(cd "$deploy_dir" && pwd -P)
if [[ -n "$label_dir" && -d "$label_dir" ]]; then
label_dir=$(cd "$label_dir" && pwd -P)
fi
if [[ -n "$requested_dir" && -n "$label_dir" && "$deploy_dir" != "$label_dir" ]]; then
echo "❌ 指定部署目录与运行中容器标签不一致:$deploy_dir != $label_dir"
return 1
fi
project_name="${COMPOSE_PROJECT_NAME:-}"
if [[ -z "$project_name" && -n "$label_project" ]]; then
project_name="$label_project"
fi
if [[ -n "$project_name" && ! "$project_name" =~ ^[a-z0-9][a-z0-9_-]*$ ]]; then
echo "❌ Compose 项目名不合法:$project_name"
return 1
fi
PANEL_DEPLOY_DIR="$deploy_dir"
PANEL_BACKEND_CONTAINER="$backend_container"
PANEL_POSTGRES_CONTAINER="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
PANEL_COMPOSE_PROJECT="$project_name"
export PANEL_DEPLOY_DIR PANEL_BACKEND_CONTAINER PANEL_POSTGRES_CONTAINER PANEL_COMPOSE_PROJECT
if [[ -n "$project_name" ]]; then
COMPOSE_PROJECT_NAME="$project_name"
export COMPOSE_PROJECT_NAME
fi
cd "$PANEL_DEPLOY_DIR"
echo "📁 部署目录:$PANEL_DEPLOY_DIR"
if [[ -n "$PANEL_COMPOSE_PROJECT" ]]; then
echo "📦 Compose 项目:$PANEL_COMPOSE_PROJECT"
fi
}
run_panel_compose() {
$DOCKER_CMD "$@"
}
validate_panel_update_environment() {
local key value backend_port frontend_port configured_db_type db_type database_url postgres_password
if [[ ! -f ".env" ]]; then
echo "❌ 部署目录缺少 .env,更新终止"
return 1
fi
if [[ ! -f "docker-compose.yml" ]]; then
echo "❌ 部署目录缺少 docker-compose.yml,更新终止"
return 1
fi
for key in JWT_SECRET BACKEND_PORT FRONTEND_PORT; do
value=$(get_env_var "$key" || true)
if [[ -z "${value//[[:space:]]/}" ]]; then
echo "❌ .env 中 $key 缺失或为空,更新终止"
return 1
fi
done
backend_port=$(get_env_var "BACKEND_PORT" || true)
frontend_port=$(get_env_var "FRONTEND_PORT" || true)
if [[ ! "$backend_port" =~ ^[0-9]+$ || "$backend_port" -lt 1 || "$backend_port" -gt 65535 ]]; then
echo "❌ BACKEND_PORT 不是有效端口:$backend_port"
return 1
fi
if [[ ! "$frontend_port" =~ ^[0-9]+$ || "$frontend_port" -lt 1 || "$frontend_port" -gt 65535 ]]; then
echo "❌ FRONTEND_PORT 不是有效端口:$frontend_port"
return 1
fi
configured_db_type=$(get_env_var "DB_TYPE" || true)
if [[ -n "$configured_db_type" && "$configured_db_type" != "sqlite" && "$configured_db_type" != "postgres" ]]; then
echo "❌ DB_TYPE 仅支持 sqlite 或 postgres:$configured_db_type"
return 1
fi
db_type=$(get_current_db_type)
if [[ "$db_type" == "postgres" ]]; then
database_url=$(get_env_var "DATABASE_URL" || true)
postgres_password=$(get_env_var "POSTGRES_PASSWORD" || true)
if [[ -z "${database_url//[[:space:]]/}" || -z "${postgres_password//[[:space:]]/}" ]]; then
echo "❌ PostgreSQL 模式要求 DATABASE_URL 和 POSTGRES_PASSWORD 均非空"
return 1
fi
fi
if ! run_panel_compose -f docker-compose.yml config -q; then
echo "❌ 当前 Compose 配置校验失败,更新终止"
return 1
fi
}
download_panel_update_compose() {
local compose_url="$1"
UPDATE_COMPOSE_CANDIDATE=$(mktemp "$PANEL_DEPLOY_DIR/.docker-compose.yml.update.XXXXXX")
if ! curl -fL -o "$UPDATE_COMPOSE_CANDIDATE" "$compose_url"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 下载最新 Compose 配置失败,更新终止"
return 1
fi
if [[ ! -s "$UPDATE_COMPOSE_CANDIDATE" ]]; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 下载的 Compose 配置为空,更新终止"
return 1
fi
}
validate_panel_update_compose() {
if ! FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" config -q; then
echo "❌ 新 Compose 配置校验失败,更新终止"
return 1
fi
}
backup_sqlite_for_update() (
local destination="$1"
local running paused="false"
cleanup_paused_backend() {
if [[ "$paused" == "true" ]]; then
docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null 2>&1 || true
fi
}
trap cleanup_paused_backend EXIT
running=$(docker inspect -f '{{.State.Running}}' "$PANEL_BACKEND_CONTAINER" 2>/dev/null || true)
if [[ "$running" == "true" ]]; then
if ! docker pause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
echo "❌ 暂停后端以创建一致性 SQLite 备份失败"
return 1
fi
paused="true"
fi
mkdir -p "$destination"
if ! docker cp "$PANEL_BACKEND_CONTAINER:/app/data/." "$destination"; then
echo "❌ SQLite 数据备份失败"
return 1
fi
if [[ "$paused" == "true" ]] && ! docker unpause "$PANEL_BACKEND_CONTAINER" >/dev/null; then
echo "❌ SQLite 备份完成,但恢复后端运行失败"
return 1
fi
paused="false"
)
backup_postgres_for_update() {
local destination="$1"
local postgres_db postgres_user
postgres_db=$(get_env_var "POSTGRES_DB" || true)
postgres_user=$(get_env_var "POSTGRES_USER" || true)
postgres_db=${postgres_db:-flux_panel}
postgres_user=${postgres_user:-flux_panel}
if ! docker exec "$PANEL_POSTGRES_CONTAINER" pg_dump -U "$postgres_user" "$postgres_db" > "$destination"; then
rm -f "$destination"
echo "❌ PostgreSQL 数据备份失败"
return 1
fi
}
create_panel_update_backup() {
local db_type="$1"
local timestamp
timestamp=$(date '+%Y%m%d-%H%M%S')
if ! (umask 077 && mkdir -p "$PANEL_DEPLOY_DIR/backups"); then
echo "❌ 创建更新备份目录失败"
return 1
fi
UPDATE_BACKUP_DIR=$(umask 077 && mktemp -d "$PANEL_DEPLOY_DIR/backups/panel-update-$timestamp.XXXXXX") || {
echo "❌ 创建更新备份目录失败"
return 1
}
if ! cp -p .env docker-compose.yml "$UPDATE_BACKUP_DIR/"; then
echo "❌ 备份面板配置失败"
return 1
fi
if [[ "$db_type" == "postgres" ]]; then
backup_postgres_for_update "$UPDATE_BACKUP_DIR/postgres.sql" || return 1
else
backup_sqlite_for_update "$UPDATE_BACKUP_DIR/sqlite" || return 1
fi
echo "💾 更新备份:$UPDATE_BACKUP_DIR"
}
pull_panel_update_images() {
local db_type="$1"
if [[ "$db_type" == "postgres" ]]; then
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend postgres
else
FLUX_VERSION="$LATEST_VERSION" run_panel_compose -f "$UPDATE_COMPOSE_CANDIDATE" pull backend frontend
fi
}
activate_panel_update_files() {
if ! chmod 0644 "$UPDATE_COMPOSE_CANDIDATE" || ! mv "$UPDATE_COMPOSE_CANDIDATE" docker-compose.yml; then
echo "❌ 替换 Compose 配置失败"
return 1
fi
UPDATE_COMPOSE_CANDIDATE=""
if ! upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"; then
echo "❌ 更新版本配置失败"
return 1
fi
}
start_panel_after_update() {
local db_type="$1"
if [[ "$db_type" == "postgres" ]]; then
run_panel_compose up -d postgres || return 1
wait_for_postgres_healthy || return 1
fi
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
wait_for_backend_healthy
}
rollback_panel_update() {
local db_type="$1"
echo "↩️ 正在恢复更新前配置..."
if ! cp -p "$UPDATE_BACKUP_DIR/.env" .env || ! cp -p "$UPDATE_BACKUP_DIR/docker-compose.yml" docker-compose.yml; then
echo "❌ 配置回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
return 1
fi
if [[ "$db_type" == "postgres" ]]; then
run_panel_compose up -d postgres || return 1
wait_for_postgres_healthy || return 1
fi
run_panel_compose up -d --force-recreate --remove-orphans backend frontend || return 1
wait_for_backend_healthy
}
wait_for_postgres_healthy() {
local pg_health
local postgres_container="${PANEL_POSTGRES_CONTAINER:-$DEFAULT_PANEL_POSTGRES_CONTAINER}"
echo "🔍 检查 PostgreSQL 服务状态..."
for i in {1..90}; do
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-postgres$"; then
pg_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo "unknown")
if docker ps --format "{{.Names}}" | grep -Fxq "$postgres_container"; then
pg_health=$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo "unknown")
if [[ "$pg_health" == "healthy" ]]; then
echo "✅ PostgreSQL 服务健康检查通过"
return 0
@@ -335,7 +614,7 @@ wait_for_postgres_healthy() {
if [ $i -eq 90 ]; then
echo "❌ PostgreSQL 启动超时(90秒)"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-postgres 2>/dev/null || echo '容器不存在')"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$postgres_container" 2>/dev/null || echo '容器不存在')"
return 1
fi
@@ -348,11 +627,12 @@ wait_for_postgres_healthy() {
wait_for_backend_healthy() {
local backend_health
local backend_container="${PANEL_BACKEND_CONTAINER:-$DEFAULT_PANEL_BACKEND_CONTAINER}"
echo "🔍 检查后端服务状态..."
for i in {1..90}; do
if docker ps --format "{{.Names}}" | grep -q "^flux-panel-backend$"; then
backend_health=$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo "unknown")
if docker ps --format "{{.Names}}" | grep -Fxq "$backend_container"; then
backend_health=$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo "unknown")
if [[ "$backend_health" == "healthy" ]]; then
echo "✅ 后端服务健康检查通过"
return 0
@@ -365,7 +645,7 @@ wait_for_backend_healthy() {
if [ $i -eq 90 ]; then
echo "❌ 后端服务启动超时(90秒)"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' flux-panel-backend 2>/dev/null || echo '容器不存在')"
echo "🔍 当前状态:$(docker inspect -f '{{.State.Health.Status}}' "$backend_container" 2>/dev/null || echo '容器不存在')"
return 1
fi
@@ -380,9 +660,8 @@ wait_for_backend_healthy() {
delete_self() {
echo ""
echo "🗑️ 操作已完成,正在清理脚本文件..."
SCRIPT_PATH="$(readlink -f "$0" 2>/dev/null || realpath "$0" 2>/dev/null || echo "$0")"
sleep 1
rm -f "$SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
rm -f "$PANEL_SCRIPT_PATH" && echo "✅ 脚本文件已删除" || echo "❌ 删除脚本文件失败"
}
@@ -486,12 +765,11 @@ EOF
# 更新功能
update_panel() {
echo "🔄 开始更新面板..."
ask_proxy_config
check_docker
resolve_panel_deployment || return 1
ask_proxy_config
validate_panel_update_environment || return 1
if [[ ! -f ".env" ]]; then
echo "⚠️ 未找到 .env,默认按 SQLite 模式更新"
fi
CURRENT_DB_TYPE=$(get_current_db_type)
echo "🗄️ 当前数据库类型:$CURRENT_DB_TYPE"
@@ -502,52 +780,45 @@ update_panel() {
}
echo "🆕 最新版本:$LATEST_VERSION"
set_compose_urls_by_version "$LATEST_VERSION"
upsert_env_var ".env" "FLUX_VERSION" "$LATEST_VERSION"
echo "🔽 下载最新配置文件..."
DOCKER_COMPOSE_URL=$(get_docker_compose_url)
echo "📡 选择配置文件:$(basename "$DOCKER_COMPOSE_URL")"
curl -L -o docker-compose.yml "$DOCKER_COMPOSE_URL"
echo "✅ 下载完成"
# 自动检测并配置 IPv6 支持
if check_ipv6_support; then
echo "🚀 系统支持 IPv6,自动启用 IPv6 配置..."
configure_docker_ipv6
download_panel_update_compose "$DOCKER_COMPOSE_URL" || return 1
if ! validate_panel_update_compose; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
return 1
fi
# 先发送 SIGTERM 信号,让应用优雅关闭
docker stop -t 30 flux-panel-backend 2>/dev/null || true
docker stop -t 10 vite-frontend 2>/dev/null || true
# 等待 WAL 文件同步
echo "⏳ 等待数据同步..."
sleep 5
# 然后再完全停止
$DOCKER_CMD down
echo "💾 备份当前配置和数据库..."
if ! create_panel_update_backup "$CURRENT_DB_TYPE"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
return 1
fi
echo "⬇️ 拉取最新镜像..."
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
$DOCKER_CMD pull backend frontend postgres
else
$DOCKER_CMD pull backend frontend
if ! pull_panel_update_images "$CURRENT_DB_TYPE"; then
rm -f "$UPDATE_COMPOSE_CANDIDATE"
UPDATE_COMPOSE_CANDIDATE=""
echo "❌ 镜像拉取失败,现有服务保持运行"
return 1
fi
echo "🚀 启动更新后的服务..."
if [[ "$CURRENT_DB_TYPE" == "postgres" ]]; then
$DOCKER_CMD up -d postgres
wait_for_postgres_healthy
$DOCKER_CMD up -d backend frontend
else
$DOCKER_CMD up -d backend frontend
if ! activate_panel_update_files; then
rollback_panel_update "$CURRENT_DB_TYPE" || true
return 1
fi
# 等待服务启动
echo "⏳ 等待服务启动..."
if ! wait_for_backend_healthy; then
echo "🛑 更新终止"
echo "🚀 原地重建更新后的服务..."
if ! start_panel_after_update "$CURRENT_DB_TYPE"; then
echo "🛑 新版本启动失败"
if rollback_panel_update "$CURRENT_DB_TYPE"; then
echo "✅ 已恢复更新前版本"
else
echo "❌ 自动回滚失败,请从 $UPDATE_BACKUP_DIR 手动恢复"
fi
return 1
fi
+262 -7
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"
@@ -272,6 +277,119 @@ EOF
local expected=$'{\n "addr": "panel\\"addr",\n "secret": "sec\\\\ret\\"1"\n}'
assert_equals "$expected" "$actual" "install_flux_agent should JSON-escape config values"
grep -Fqx 'Environment=GODEBUG=disablethp=1' "$FLUX_AGENT_SYSTEMD_SERVICE_FILE" || \
fail "systemd service should disable transparent huge pages for the Go heap"
)
test_install_script_bootstraps_bash_for_alpine() (
set -euo pipefail
local shebang
shebang=$(head -n 1 "$ROOT_DIR/install.sh")
assert_equals "#!/bin/sh" "$shebang" "install.sh should start with Alpine's default shell"
grep -Fq 'apk add --no-cache bash' "$ROOT_DIR/install.sh" || \
fail "install.sh should bootstrap Bash through apk on Alpine"
)
test_install_flux_agent_uses_openrc() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
local temp_root
temp_root=$(mktemp -d)
INSTALL_DIR="$temp_root/flux_agent"
FLUX_AGENT_OPENRC_SERVICE_FILE="$temp_root/init.d/flux_agent"
SERVICE_MANAGER="openrc"
SERVER_ADDR="panel.example.com:443"
SECRET="secret"
DOWNLOAD_URL="https://example.com/gost"
local rc_service_calls=""
local rc_update_calls=""
ask_proxy_config() { :; }
ensure_download_url_initialized() { :; }
get_config_params() { :; }
check_and_install_tcpkill() { :; }
cleanup_legacy_gost_installation() { :; }
curl() {
local output=""
while [[ $# -gt 0 ]]; do
if [[ "$1" == "-o" ]]; then
output="$2"
shift 2
continue
fi
shift
done
cat > "$output" <<'EOF'
#!/bin/sh
echo "new version"
EOF
chmod +x "$output"
}
rc-service() {
rc_service_calls+=$'\n'"$*"
if [[ "$2" == "status" ]]; then
echo "status: started"
fi
return 0
}
rc-update() {
rc_update_calls+=$'\n'"$*"
return 0
}
install_flux_agent >/dev/null
[[ -x "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]] || fail "OpenRC service file should be executable"
grep -Fq '#!/sbin/openrc-run' "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
fail "OpenRC service should use openrc-run"
grep -Fq "command=\"$INSTALL_DIR/flux_agent\"" "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
fail "OpenRC service should launch the installed flux_agent binary"
grep -Fq 'command_background="yes"' "$FLUX_AGENT_OPENRC_SERVICE_FILE" || \
fail "OpenRC service should run flux_agent in the background"
if command -v openrc-run >/dev/null 2>&1; then
"$FLUX_AGENT_OPENRC_SERVICE_FILE" describe >/dev/null 2>&1
fi
[[ "$rc_update_calls" == *"add flux_agent default"* ]] || \
fail "OpenRC install should enable flux_agent in the default runlevel"
[[ "$rc_service_calls" == *"start"* ]] || fail "OpenRC install should start flux_agent"
[[ "$rc_service_calls" == *"status"* ]] || fail "OpenRC install should verify flux_agent status"
)
test_remove_flux_agent_service_uses_openrc() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
local temp_root
temp_root=$(mktemp -d)
SERVICE_MANAGER="openrc"
FLUX_AGENT_OPENRC_SERVICE_FILE="$temp_root/init.d/flux_agent"
mkdir -p "$(dirname "$FLUX_AGENT_OPENRC_SERVICE_FILE")"
: > "$FLUX_AGENT_OPENRC_SERVICE_FILE"
local rc_service_calls=""
local rc_update_calls=""
rc-service() {
rc_service_calls+=$'\n'"$*"
return 0
}
rc-update() {
rc_update_calls+=$'\n'"$*"
return 0
}
stop_flux_agent_service
disable_flux_agent_service
remove_flux_agent_service
[[ "$rc_service_calls" == *"stop"* ]] || fail "OpenRC uninstall should stop flux_agent"
[[ "$rc_update_calls" == *"del flux_agent default"* ]] || \
fail "OpenRC uninstall should remove flux_agent from the default runlevel"
[[ ! -e "$FLUX_AGENT_OPENRC_SERVICE_FILE" ]] || fail "OpenRC uninstall should remove its service file"
)
test_cleanup_legacy_gost_installation_removes_service_and_binary() (
@@ -283,6 +401,7 @@ test_cleanup_legacy_gost_installation_removes_service_and_binary() (
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
SERVICE_MANAGER="systemd"
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
[Unit]
Description=Gost Proxy Service
@@ -326,6 +445,7 @@ test_cleanup_legacy_gost_installation_preserves_unrelated_gost() (
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
SERVICE_MANAGER="systemd"
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<'EOF'
[Unit]
Description=Unrelated Gost Service
@@ -353,6 +473,35 @@ EOF
[[ "$systemctl_calls" != *"disable gost"* ]] || fail "cleanup_legacy_gost_installation should not disable unrelated gost services"
)
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
LEGACY_GOST_SERVICE_FILE_ETC=$(mktemp)
LEGACY_GOST_SERVICE_FILE_LIB=$(mktemp -u)
LEGACY_GOST_SERVICE_FILE_USR_LIB=$(mktemp -u)
LEGACY_GOST_CONFIG_DIR=$(mktemp -d)
SERVICE_MANAGER="openrc"
cat > "$LEGACY_GOST_SERVICE_FILE_ETC" <<EOF
[Unit]
WorkingDirectory=$LEGACY_GOST_CONFIG_DIR
ExecStart=$LEGACY_GOST_CONFIG_DIR/gost
EOF
: > "$LEGACY_GOST_CONFIG_DIR/config.json"
: > "$LEGACY_GOST_CONFIG_DIR/gost.json"
local systemctl_calls=""
systemctl() {
systemctl_calls+=$'\n'"$*"
return 1
}
cleanup_legacy_gost_installation >/dev/null
[[ -z "$systemctl_calls" ]] || fail "OpenRC cleanup should not invoke systemctl"
[[ ! -e "$LEGACY_GOST_SERVICE_FILE_ETC" ]] || fail "OpenRC cleanup should still remove the legacy service file"
)
test_install_script_accepts_proxy_url_env_without_prompt() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/install.sh"
@@ -409,6 +558,7 @@ test_update_panel_asks_for_proxy_config() (
load_script_without_main "$ROOT_DIR/panel_install.sh"
local ask_called="0"
local calls=""
ask_proxy_config() {
ask_called="1"
@@ -421,6 +571,15 @@ test_update_panel_asks_for_proxy_config() (
DOCKER_CMD="true"
}
resolve_panel_deployment() {
calls+=" resolve"
PANEL_DEPLOY_DIR=$(mktemp -d)
}
validate_panel_update_environment() {
calls+=" validate-current"
}
get_current_db_type() {
echo "sqlite"
}
@@ -429,13 +588,33 @@ test_update_panel_asks_for_proxy_config() (
echo "v-test"
}
upsert_env_var() { :; }
download_panel_update_compose() {
calls+=" download"
UPDATE_COMPOSE_CANDIDATE="candidate"
}
validate_panel_update_compose() {
calls+=" validate-new"
}
create_panel_update_backup() {
calls+=" backup"
UPDATE_BACKUP_DIR="backup"
}
pull_panel_update_images() {
calls+=" pull"
}
activate_panel_update_files() {
calls+=" activate"
}
start_panel_after_update() {
calls+=" start"
}
check_ipv6_support() { return 1; }
configure_docker_ipv6() { :; }
docker() { return 0; }
wait_for_backend_healthy() { return 0; }
sleep() { :; }
curl() { :; }
update_panel >/dev/null
@@ -444,6 +623,75 @@ test_update_panel_asks_for_proxy_config() (
"https://github.com/${REPO}/releases/download/v-test/docker-compose-v4.yml" \
"$DOCKER_COMPOSEV4_URL" \
"update_panel should honor the prompted proxy choice"
assert_equals \
" resolve validate-current download validate-new backup pull activate start" \
"$calls" \
"update_panel should validate, back up, and pull before activating the update"
)
test_resolve_panel_deployment_uses_container_labels() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/panel_install.sh"
local expected_dir
expected_dir=$(mktemp -d)
expected_dir=$(cd "$expected_dir" && pwd -P)
printf 'JWT_SECRET=test\nBACKEND_PORT=6365\nFRONTEND_PORT=6366\nDB_TYPE=sqlite\n' > "$expected_dir/.env"
printf 'services: {}\n' > "$expected_dir/docker-compose.yml"
unset PANEL_DEPLOY_DIR COMPOSE_PROJECT_NAME
docker() {
if [[ "$1" == "inspect" && "$2" == "-f" ]]; then
case "$3" in
*working_dir*) printf '%s\n' "$expected_dir" ;;
*) printf '%s\n' "flvx-panel" ;;
esac
fi
}
resolve_panel_deployment >/dev/null
assert_equals "$expected_dir" "$PANEL_DEPLOY_DIR" "update should discover the Compose working directory from container labels"
assert_equals "flvx-panel" "$COMPOSE_PROJECT_NAME" "update should preserve the existing Compose project name"
assert_equals "$expected_dir" "$(pwd -P)" "update should run from the discovered deployment directory"
)
test_validate_panel_update_environment_rejects_empty_required_values() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/panel_install.sh"
local deploy_dir rc="0"
deploy_dir=$(mktemp -d)
printf 'JWT_SECRET=\nBACKEND_PORT=6365\nFRONTEND_PORT=6366\nDB_TYPE=sqlite\n' > "$deploy_dir/.env"
printf 'services: {}\n' > "$deploy_dir/docker-compose.yml"
cd "$deploy_dir"
DOCKER_CMD="true"
validate_panel_update_environment >/dev/null || rc="$?"
assert_equals "1" "$rc" "update should reject an empty JWT secret before touching the running service"
)
test_backup_sqlite_for_update_pauses_copies_and_unpauses() (
set -euo pipefail
load_script_without_main "$ROOT_DIR/panel_install.sh"
local destination calls log_path
destination=$(mktemp -d)/sqlite
log_path=$(mktemp)
PANEL_BACKEND_CONTAINER="flux-panel-backend"
docker() {
if [[ "$1" == "inspect" ]]; then
printf 'true\n'
return 0
fi
printf ' %s' "$1" >> "$log_path"
}
backup_sqlite_for_update "$destination" >/dev/null
calls=$(cat "$log_path")
assert_equals " pause cp unpause" "$calls" "SQLite backup should pause writes, copy the database volume, then resume the backend"
)
test_panel_install_script_accepts_proxy_url_env_without_prompt() (
@@ -504,14 +752,21 @@ test_update_flux_agent_skips_proxy_prompt_when_not_installed
test_install_flux_agent_preserves_legacy_gost_when_download_fails
test_update_flux_agent_preserves_legacy_gost_when_download_fails
test_install_flux_agent_writes_json_safe_config
test_install_script_bootstraps_bash_for_alpine
test_install_flux_agent_uses_openrc
test_remove_flux_agent_service_uses_openrc
test_cleanup_legacy_gost_installation_removes_service_and_binary
test_cleanup_legacy_gost_installation_preserves_unrelated_gost
test_cleanup_legacy_gost_installation_skips_systemd_on_openrc
test_install_script_accepts_proxy_url_env_without_prompt
test_panel_install_script_can_disable_proxy
test_panel_install_script_recomputes_compose_urls_after_prompt
test_update_panel_asks_for_proxy_config
test_resolve_panel_deployment_uses_container_labels
test_validate_panel_update_environment_rejects_empty_required_values
test_backup_sqlite_for_update_pauses_copies_and_unpauses
test_panel_install_script_uses_default_proxy
test_panel_install_script_accepts_proxy_url_env_without_prompt
test_panel_install_script_defaults_proxy_on_eof
echo "install script proxy tests passed"
echo "install script tests passed"
+3 -3
View File
@@ -3,8 +3,8 @@ FROM node:22-alpine AS builder
WORKDIR /app
COPY package.json pnpm-lock.yaml* ./
RUN corepack prepare pnpm@10 --activate && corepack enable pnpm && pnpm install --frozen-lockfile
COPY package.json pnpm-lock.yaml pnpm-workspace.yaml ./
RUN corepack enable pnpm && pnpm install --frozen-lockfile
COPY . .
RUN pnpm run build
@@ -18,4 +18,4 @@ COPY --from=builder /app/dist /usr/share/nginx/html
EXPOSE 80
CMD ["nginx", "-g", "daemon off;"]
CMD ["nginx", "-g", "daemon off;"]
+1 -7
View File
@@ -2,6 +2,7 @@
"name": "flvx",
"private": true,
"version": "0.0.0",
"packageManager": "pnpm@10.28.1",
"type": "module",
"scripts": {
"dev": "vite",
@@ -81,12 +82,5 @@
"vite": "npm:rolldown-vite@^7.3.1",
"vite-plugin-pwa": "^1.1.0",
"vite-tsconfig-paths": "^6.0.5"
},
"pnpm": {
"overrides": {
"@babel/plugin-transform-modules-systemjs": "7.29.4",
"fast-uri": "3.1.2",
"serialize-javascript": "7.0.5"
}
}
}
+8
View File
@@ -1,2 +1,10 @@
packages:
- '.'
overrides:
'@babel/plugin-transform-modules-systemjs': 7.29.4
fast-uri: 3.1.2
serialize-javascript: 7.0.5
allowBuilds:
'@tailwindcss/oxide': true
+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) =>
+14
View File
@@ -595,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,
+74 -22
View File
@@ -14,10 +14,14 @@ const PUBLIC_BRAND_CONFIG_KEYS = [
"app_logo",
"app_favicon",
"app_bg_image",
"is_commercial",
"hide_footer_brand",
] as const;
const SENSITIVE_CONFIG_KEYS = new Set([
"jwt_secret",
"license_key",
"license_machine_id",
"machine_fingerprint",
"cloudflare_secret_key",
]);
const GITHUB_REPO =
@@ -49,6 +53,30 @@ const readCachedConfigs = (keys: readonly string[]) => {
return { cachedConfigs, hasCachedData };
};
const readAllCachedSafeConfigs = () => {
const cachedConfigs: Record<string, string> = {};
Object.keys(localStorage).forEach((storageKey) => {
if (!storageKey.startsWith(CACHE_PREFIX)) {
return;
}
const key = storageKey.slice(CACHE_PREFIX.length).trim().toLowerCase();
if (!key || SENSITIVE_CONFIG_KEYS.has(key)) {
return;
}
const value = localStorage.getItem(storageKey);
if (value !== null) {
cachedConfigs[key] = value;
}
});
return cachedConfigs;
};
const fetchPublicBrandConfigs = async (): Promise<Record<string, string>> => {
const publicConfigMap: Record<string, string> = {};
@@ -106,15 +134,15 @@ const getInitialConfig = () => {
if (cachedAppName) {
return {
name: cachedAppName,
name: isCommercial ? cachedAppName : "FLVX",
version: VERSION,
app_version: APP_VERSION,
github_repo: GITHUB_REPO,
app_logo: cachedAppLogo,
app_favicon: cachedAppFavicon,
app_logo: isCommercial ? cachedAppLogo : "",
app_favicon: isCommercial ? cachedAppFavicon : "",
app_bg_image: cachedAppBgImage,
is_commercial: isCommercial,
hide_footer_brand: hideFooterBrand,
hide_footer_brand: isCommercial && hideFooterBrand,
};
}
@@ -123,11 +151,11 @@ const getInitialConfig = () => {
version: VERSION,
app_version: APP_VERSION,
github_repo: GITHUB_REPO,
app_logo: cachedAppLogo,
app_favicon: cachedAppFavicon,
app_logo: isCommercial ? cachedAppLogo : "",
app_favicon: isCommercial ? cachedAppFavicon : "",
app_bg_image: cachedAppBgImage,
is_commercial: isCommercial,
hide_footer_brand: hideFooterBrand,
hide_footer_brand: isCommercial && hideFooterBrand,
};
};
@@ -206,20 +234,24 @@ export const getCachedConfig = async (key: string): Promise<string | null> => {
// 获取所有配置(优先从缓存)
export const getCachedConfigs = async (): Promise<Record<string, string>> => {
const { cachedConfigs, hasCachedData } = readCachedConfigs(
PUBLIC_BRAND_CONFIG_KEYS,
);
const {
cachedConfigs: publicCachedConfigs,
hasCachedData: hasPublicCachedData,
} = readCachedConfigs(PUBLIC_BRAND_CONFIG_KEYS);
if (!isLoggedIn()) {
const publicConfigs = await fetchPublicBrandConfigs();
if (Object.keys(publicConfigs).length > 0) {
return { ...cachedConfigs, ...publicConfigs };
return { ...publicCachedConfigs, ...publicConfigs };
}
return cachedConfigs;
return publicCachedConfigs;
}
const cachedConfigs = readAllCachedSafeConfigs();
const hasCachedData = Object.keys(cachedConfigs).length > 0;
// 从API获取最新配置
try {
const response = await getConfigs();
@@ -249,14 +281,20 @@ export const getCachedConfigs = async (): Promise<Record<string, string>> => {
return cachedConfigs;
}
return await fetchPublicBrandConfigs();
const publicConfigs = await fetchPublicBrandConfigs();
return { ...publicCachedConfigs, ...publicConfigs };
} catch {
// API失败时返回缓存的数据
if (hasCachedData) {
return cachedConfigs;
}
return await fetchPublicBrandConfigs();
const publicConfigs = await fetchPublicBrandConfigs();
return hasPublicCachedData
? { ...publicCachedConfigs, ...publicConfigs }
: publicConfigs;
}
};
@@ -345,6 +383,12 @@ export const updateSiteConfig = async (configMap?: Record<string, string>) => {
"app_bg_image",
);
const resolvedCommercial = Object.prototype.hasOwnProperty.call(
resolvedConfigMap,
"is_commercial",
)
? resolvedConfigMap.is_commercial === "true"
: siteConfig.is_commercial;
const appName = hasAppName
? String(resolvedConfigMap.app_name || "").trim()
: siteConfig.name;
@@ -358,15 +402,23 @@ export const updateSiteConfig = async (configMap?: Record<string, string>) => {
? String(resolvedConfigMap.app_bg_image || "").trim()
: (siteConfig.app_bg_image || "").trim();
if (appName && appName !== siteConfig.name) {
siteConfig.name = appName;
}
siteConfig.app_logo = appLogo;
siteConfig.app_favicon = appFavicon;
siteConfig.name = resolvedCommercial && appName ? appName : "FLVX";
siteConfig.app_logo = resolvedCommercial ? appLogo : "";
siteConfig.app_favicon = resolvedCommercial ? appFavicon : "";
siteConfig.app_bg_image = appBgImage;
siteConfig.is_commercial = resolvedConfigMap.is_commercial === "true";
siteConfig.hide_footer_brand = resolvedConfigMap.hide_footer_brand === "true";
if (
Object.prototype.hasOwnProperty.call(resolvedConfigMap, "is_commercial")
) {
siteConfig.is_commercial = resolvedCommercial;
}
if (
Object.prototype.hasOwnProperty.call(resolvedConfigMap, "hide_footer_brand")
) {
siteConfig.hide_footer_brand =
resolvedCommercial && resolvedConfigMap.hide_footer_brand === "true";
} else if (!resolvedCommercial) {
siteConfig.hide_footer_brand = false;
}
if (typeof document !== "undefined") {
document.title = siteConfig.name;
@@ -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} 秒`;
+75 -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",
@@ -247,6 +266,9 @@ const getInitialConfigs = (): Record<string, string> => {
"github_proxy_enabled",
"github_proxy_url",
"allow_local_remote_addr",
"is_commercial",
"license_expiry",
"hide_footer_brand",
];
const initialConfigs: Record<string, string> = {};
@@ -622,6 +644,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 +702,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 +1120,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 +1146,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 (
+180 -4
View File
@@ -69,6 +69,7 @@ import {
getNodeList,
pauseForwardService,
resumeForwardService,
resetForwardFlow,
diagnoseForward,
updateForwardOrder,
getConfigByName,
@@ -230,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 => {
@@ -764,6 +765,7 @@ const SortableTableRow = ({
handleEdit,
handleDelete,
handleDiagnose,
handleResetFlow,
showAddressModal,
formatFlow,
}: any) => {
@@ -919,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"
@@ -958,6 +983,7 @@ const SortableCompactTableRow = ({
handleEdit,
handleDelete,
handleDiagnose,
handleResetFlow,
showAddressModal,
hasMultipleAddresses,
formatFlow,
@@ -1144,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"
@@ -1289,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] =
@@ -2171,7 +2225,8 @@ export default function ForwardPage() {
maxConn: forward.maxConn ?? 0,
proxyProtocol: forward.proxyProtocol ?? 0,
proxyProtocolReceive: forward.proxyProtocolReceive ?? 0,
proxyProtocolSend: forward.proxyProtocolSend ?? forward.proxyProtocol ?? 0,
proxyProtocolSend:
forward.proxyProtocolSend ?? forward.proxyProtocol ?? 0,
});
setErrors({});
setModalOpen(true);
@@ -2183,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;
@@ -3957,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"
@@ -4000,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"
@@ -4302,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>
@@ -4321,6 +4442,7 @@ export default function ForwardPage() {
handleDelete={handleDelete}
handleDiagnose={handleDiagnose}
handleEdit={handleEdit}
handleResetFlow={handleResetFlow}
handleServiceToggle={handleServiceToggle}
hasMultipleAddresses={hasMultipleAddresses}
selectMode={selectMode}
@@ -4616,6 +4738,7 @@ export default function ForwardPage() {
handleDelete={handleDelete}
handleDiagnose={handleDiagnose}
handleEdit={handleEdit}
handleResetFlow={handleResetFlow}
handleServiceToggle={
handleServiceToggle
}
@@ -5172,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) {