mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
Compare commits
7 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 799bb66fe5 | |||
| 3f374df724 | |||
| 9b923a2d0b | |||
| d9f28f53c7 | |||
| 1d08a1ccfc | |||
| db25ba2cbe | |||
| a498067261 |
@@ -0,0 +1,171 @@
|
||||
# Allow Local Remote Address Implementation Plan
|
||||
|
||||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||
|
||||
**Goal:** Add a global settings toggle that allows non-admin forward rules to target local/private addresses when explicitly enabled.
|
||||
|
||||
**Architecture:** Keep the existing remote-address safety validator as the default path for non-admin rule changes, but gate its use behind a single backend config lookup in forward create/update handlers. Surface the toggle through the existing `vite_config` settings page and prove behavior with backend contract tests first.
|
||||
|
||||
**Tech Stack:** Go `net/http` + GORM backend, React + TypeScript frontend settings page, Go contract tests.
|
||||
|
||||
---
|
||||
|
||||
### Task 1: Backend Contract Coverage
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/tests/contract/forward_contract_test.go`
|
||||
|
||||
- [ ] **Step 1: Write the failing tests**
|
||||
|
||||
Add contract tests that prove the desired behavior:
|
||||
|
||||
```go
|
||||
t.Run("local remote address is rejected when toggle is off", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "deny-local-remote",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, _ := json.Marshal(createPayload)
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
|
||||
var out response.R
|
||||
_ = json.NewDecoder(createRes.Body).Decode(&out)
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected local remote address to be rejected when toggle is off")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("local remote address is allowed when toggle is on", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, "allow_local_remote_addr", "1", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("enable allow_local_remote_addr: %v", err)
|
||||
}
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "allow-local-remote",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, _ := json.Marshal(createPayload)
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
})
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Run tests to verify they fail**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: FAIL because backend still rejects local/private addresses unconditionally.
|
||||
|
||||
- [ ] **Step 3: Commit**
|
||||
|
||||
Do not commit yet; combine with Task 2 after implementation passes.
|
||||
|
||||
### Task 2: Backend Toggle Implementation
|
||||
|
||||
**Files:**
|
||||
- Modify: `go-backend/internal/http/handler/mutations.go`
|
||||
|
||||
- [ ] **Step 1: Add a tiny config helper**
|
||||
|
||||
Add a helper near other handler helpers:
|
||||
|
||||
```go
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
cfg, err := h.repo.GetConfigByName("allow_local_remote_addr")
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(cfg.Value) == "1"
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 2: Gate create/update validation behind the helper**
|
||||
|
||||
Replace the unconditional checks with:
|
||||
|
||||
```go
|
||||
if !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
- [ ] **Step 3: Run contract tests to verify they pass**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 4: Run full backend tests**
|
||||
|
||||
Run: `go test ./...`
|
||||
Expected: PASS
|
||||
|
||||
### Task 3: Settings Page Toggle
|
||||
|
||||
**Files:**
|
||||
- Modify: `vite-frontend/src/pages/config.tsx`
|
||||
|
||||
- [ ] **Step 1: Add the config item to the settings schema**
|
||||
|
||||
Add a switch-style item for `allow_local_remote_addr` with warning copy about reduced safety.
|
||||
|
||||
- [ ] **Step 2: Ensure the key is included in config loading/saving paths**
|
||||
|
||||
Add `allow_local_remote_addr` anywhere the page enumerates config keys or groups persisted config values.
|
||||
|
||||
- [ ] **Step 3: Run frontend build**
|
||||
|
||||
Run: `pnpm run build`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 4: Run frontend lint**
|
||||
|
||||
Run: `pnpm run lint`
|
||||
Expected: 0 errors; existing warnings may remain.
|
||||
|
||||
### Task 4: Final Verification
|
||||
|
||||
**Files:**
|
||||
- Verify only
|
||||
|
||||
- [ ] **Step 1: Re-run backend contracts for the toggle**
|
||||
|
||||
Run: `go test ./tests/contract/... -run 'TestForwardContracts|local remote address'`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 2: Re-run full backend tests**
|
||||
|
||||
Run: `go test ./...`
|
||||
Expected: PASS
|
||||
|
||||
- [ ] **Step 3: Re-run frontend build/lint**
|
||||
|
||||
Run: `pnpm run build && pnpm run lint`
|
||||
Expected: Build passes, lint has no errors.
|
||||
|
||||
- [ ] **Step 4: Commit**
|
||||
|
||||
```bash
|
||||
git add go-backend/internal/http/handler/mutations.go go-backend/tests/contract/forward_contract_test.go vite-frontend/src/pages/config.tsx docs/superpowers/specs/2026-04-26-allow-local-remote-addr-design.md docs/superpowers/plans/2026-04-26-allow-local-remote-addr.md
|
||||
git commit -m "feat: add allow-local-remote-address toggle"
|
||||
```
|
||||
@@ -90,7 +90,7 @@
|
||||
| 表名 | Model | 特殊处理 |
|
||||
|------|-------|----------|
|
||||
| `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) |
|
||||
| `forward` | `Forward` | |
|
||||
| `forward` | `Forward` | 增加 `proxy_protocol` 字段 |
|
||||
| `forward_port` | `ForwardPort` | |
|
||||
| `node` | `Node` | |
|
||||
| `speed_limit` | `SpeedLimit` | |
|
||||
|
||||
@@ -1705,6 +1705,12 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
if cLimiterName != "" {
|
||||
service["climiter"] = cLimiterName
|
||||
}
|
||||
if forward.ProxyProtocol > 0 {
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
service["metadata"].(map[string]interface{})["proxyProtocol"] = forward.ProxyProtocol
|
||||
}
|
||||
if protocol == "udp" {
|
||||
listenerMetadata := map[string]interface{}{
|
||||
"keepAlive": true,
|
||||
@@ -1716,7 +1722,10 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
|
||||
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
|
||||
}
|
||||
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
|
||||
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
|
||||
if service["metadata"] == nil {
|
||||
service["metadata"] = map[string]interface{}{}
|
||||
}
|
||||
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
|
||||
}
|
||||
if limiterID != nil && *limiterID > 0 {
|
||||
service["limiter"] = strconv.FormatInt(*limiterID, 10)
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
func TestBuildForwardServiceConfigsPreservesProxyProtocolWithInterfaceMetadata(t *testing.T) {
|
||||
forward := &forwardRecord{
|
||||
ID: 1,
|
||||
UserID: 2,
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
ProxyProtocol: 2,
|
||||
}
|
||||
tunnel := &tunnelRecord{Type: 1}
|
||||
node := &nodeRecord{
|
||||
InterfaceName: "eth0",
|
||||
TCPListenAddr: "0.0.0.0",
|
||||
UDPListenAddr: "0.0.0.0",
|
||||
}
|
||||
|
||||
services := buildForwardServiceConfigs("1_2_3", forward, tunnel, node, 4001, "", nil, "")
|
||||
if len(services) != 2 {
|
||||
t.Fatalf("expected 2 services, got %d", len(services))
|
||||
}
|
||||
|
||||
for _, service := range services {
|
||||
metadata, ok := service["metadata"].(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected metadata map, got %T", service["metadata"])
|
||||
}
|
||||
if metadata["interface"] != "eth0" {
|
||||
t.Fatalf("expected interface metadata eth0, got %v", metadata["interface"])
|
||||
}
|
||||
if metadata["proxyProtocol"] != 2 {
|
||||
t.Fatalf("expected proxyProtocol 2, got %v", metadata["proxyProtocol"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRollbackForwardMutationRestoresProxyProtocol(t *testing.T) {
|
||||
r, err := repo.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Create(&model.Forward{
|
||||
UserID: 2,
|
||||
UserName: "rollback-user",
|
||||
Name: "rollback-forward",
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "9.9.9.9:443",
|
||||
Strategy: "fifo",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
ProxyProtocol: 2,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create forward: %v", err)
|
||||
}
|
||||
|
||||
forwardID := mustLastInsertID(t, r, "rollback-forward")
|
||||
if err := r.DB().Model(&model.Forward{}).Where("id = ?", forwardID).Updates(map[string]interface{}{
|
||||
"name": "changed-forward",
|
||||
"proxy_protocol": 0,
|
||||
"updated_time": now + 1,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("mutate forward: %v", err)
|
||||
}
|
||||
|
||||
h := &Handler{repo: r}
|
||||
h.rollbackForwardMutation(&forwardRecord{
|
||||
ID: forwardID,
|
||||
UserID: 2,
|
||||
UserName: "rollback-user",
|
||||
Name: "rollback-forward",
|
||||
TunnelID: 3,
|
||||
RemoteAddr: "9.9.9.9:443",
|
||||
Strategy: "fifo",
|
||||
Status: 1,
|
||||
ProxyProtocol: 2,
|
||||
}, nil)
|
||||
|
||||
var proxyProtocol int
|
||||
if err := r.DB().Raw("SELECT proxy_protocol FROM forward WHERE id = ?", forwardID).Row().Scan(&proxyProtocol); err != nil {
|
||||
t.Fatalf("query proxy_protocol: %v", err)
|
||||
}
|
||||
if proxyProtocol != 2 {
|
||||
t.Fatalf("expected proxyProtocol restored to 2, got %d", proxyProtocol)
|
||||
}
|
||||
}
|
||||
@@ -50,6 +50,7 @@ type Handler struct {
|
||||
}
|
||||
|
||||
const monitorTunnelQualityEnabledConfigKey = "monitor_tunnel_quality_enabled"
|
||||
const allowLocalRemoteAddrConfigKey = "allow_local_remote_addr"
|
||||
|
||||
type loginRequest struct {
|
||||
Username string `json:"username"`
|
||||
@@ -1041,6 +1042,19 @@ func (h *Handler) isTunnelQualityMonitoringEnabled() bool {
|
||||
return strings.TrimSpace(strings.ToLower(cfg.Value)) != "false"
|
||||
}
|
||||
|
||||
func (h *Handler) allowLocalRemoteAddr() bool {
|
||||
if h == nil || h.repo == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
cfg, err := h.repo.GetConfigByName(allowLocalRemoteAddrConfigKey)
|
||||
if err != nil || cfg == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
return strings.TrimSpace(strings.ToLower(cfg.Value)) == "true"
|
||||
}
|
||||
|
||||
func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
|
||||
@@ -1726,11 +1726,13 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
|
||||
return
|
||||
}
|
||||
if roleID != 0 {
|
||||
if roleID != 0 && !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络"))
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
if roleID != 0 {
|
||||
if speedIDVal, ok := req["speedId"]; ok && speedIDVal != nil {
|
||||
response.WriteJSON(w, response.Err(-1, "普通用户无法设置限速规则"))
|
||||
return
|
||||
@@ -1780,7 +1782,9 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
userName = "user"
|
||||
}
|
||||
maxConn := asInt(req["maxConn"], 0)
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn)
|
||||
proxyProtocol := asInt(req["proxyProtocol"], 0)
|
||||
|
||||
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, inIp, nullableInt(speedID), maxConn, proxyProtocol)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1851,9 +1855,9 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if remoteAddr == "" {
|
||||
remoteAddr = forward.RemoteAddr
|
||||
}
|
||||
if actorRole != 0 {
|
||||
if actorRole != 0 && !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络"))
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -1932,8 +1936,9 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
maxConn := asInt(req["maxConn"], forward.MaxConn)
|
||||
proxyProtocol := asInt(req["proxyProtocol"], forward.ProxyProtocol)
|
||||
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn); err != nil {
|
||||
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn, proxyProtocol); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -4085,7 +4090,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
|
||||
h.repo.RollbackForwardFields(
|
||||
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
|
||||
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
|
||||
oldForward.SpeedID, oldForward.MaxConn,
|
||||
oldForward.SpeedID, oldForward.MaxConn, oldForward.ProxyProtocol,
|
||||
time.Now().UnixMilli(),
|
||||
)
|
||||
|
||||
|
||||
@@ -11,11 +11,37 @@ var DisableSafeRemoteAddrCheckForTesting = false
|
||||
|
||||
// IsSafeRemoteAddr checks if a given address is safe to connect to (prevents SSRF/Open Proxy).
|
||||
// It resolves domains to IPs to prevent DNS rebinding attacks pointing to internal networks.
|
||||
// Supports multiple addresses separated by commas or newlines (one per line).
|
||||
func IsSafeRemoteAddr(addr string) error {
|
||||
if DisableSafeRemoteAddrCheckForTesting {
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, part := range splitRemoteParts(addr) {
|
||||
if err := checkSingleRemoteAddr(part); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// splitRemoteParts splits a multi-address string by commas and newlines.
|
||||
func splitRemoteParts(addr string) []string {
|
||||
addr = strings.ReplaceAll(addr, "\n", ",")
|
||||
addr = strings.ReplaceAll(addr, "\r", ",")
|
||||
parts := strings.Split(addr, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, part := range parts {
|
||||
part = strings.TrimSpace(part)
|
||||
if part != "" {
|
||||
out = append(out, part)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// checkSingleRemoteAddr validates a single address.
|
||||
func checkSingleRemoteAddr(addr string) error {
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "missing port in address") {
|
||||
@@ -27,12 +53,12 @@ func IsSafeRemoteAddr(addr string) error {
|
||||
|
||||
ips, err := net.LookupIP(host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("could not resolve address: %v", err)
|
||||
return fmt.Errorf("could not resolve address %q: %v", addr, err)
|
||||
}
|
||||
|
||||
for _, ip := range ips {
|
||||
if ip.IsLoopback() || ip.IsPrivate() {
|
||||
return fmt.Errorf("address resolves to internal IP: %s", ip.String())
|
||||
return fmt.Errorf("address %q resolves to internal IP: %s", addr, ip.String())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -13,6 +13,13 @@ import (
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
// failedForward tracks a forward that failed redeployment, for retry.
|
||||
type failedForward struct {
|
||||
id int64
|
||||
forward *forwardRecord
|
||||
err error
|
||||
}
|
||||
|
||||
const (
|
||||
githubRepo = "Sagit-chu/flvx"
|
||||
githubAPIBase = "https://api.github.com"
|
||||
@@ -390,6 +397,9 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
|
||||
|
||||
func (h *Handler) onNodeOnline(nodeID int64) {
|
||||
h.consumeNodePendingUpgradeRedeploy(nodeID)
|
||||
// Always redeploy rules on reconnection, not just for pending upgrade nodes.
|
||||
// This handles cases where the node restarted and lost its in-memory config
|
||||
// before persistence had time to flush, or if the panel also restarted.
|
||||
h.redeployNodeRuntimeAfterUpgrade(nodeID)
|
||||
}
|
||||
|
||||
@@ -405,6 +415,7 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
return
|
||||
}
|
||||
|
||||
// First pass: deploy everything
|
||||
tunnelFailed := make(map[int64]struct{})
|
||||
for _, tunnelID := range tunnelIDs {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
@@ -413,6 +424,9 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
// Collect forwards that failed independently (not skipped due to tunnel failure)
|
||||
var failedForwards []failedForward
|
||||
|
||||
for _, forwardID := range forwardIDs {
|
||||
forward, getErr := h.getForwardRecord(forwardID)
|
||||
if getErr != nil || forward == nil {
|
||||
@@ -422,7 +436,86 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
failedForwards = append(failedForwards, failedForward{id: forwardID, forward: forward, err: err})
|
||||
fmt.Printf("post-upgrade redeploy: forward %d failed on node %d: %v\n", forwardID, nodeID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Retry failed items with exponential backoff (max 3 attempts)
|
||||
h.retryFailedRedeploys(nodeID, tunnelFailed, failedForwards)
|
||||
}
|
||||
|
||||
// isRetryableError returns true if the error looks transient and worth retrying.
|
||||
func isRetryableError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
// Skip non-retryable errors: not-found, already-exists, validation errors
|
||||
if strings.Contains(msg, "not found") || strings.Contains(msg, "不存在") {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(msg, "already exists") || strings.Contains(msg, "已存在") {
|
||||
return false
|
||||
}
|
||||
// Everything else (timeout, connection lost, port in use, etc.) is retryable
|
||||
return true
|
||||
}
|
||||
|
||||
// retryFailedRedeploys retries failed tunnels and forwards with exponential backoff.
|
||||
func (h *Handler) retryFailedRedeploys(nodeID int64, tunnelFailed map[int64]struct{}, failedForwards []failedForward) {
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
const maxRetries = 3
|
||||
baseDelay := time.Second
|
||||
|
||||
for attempt := 1; attempt <= maxRetries; attempt++ {
|
||||
delay := baseDelay * time.Duration(1<<uint(attempt-1)) // 1s, 2s, 4s
|
||||
time.Sleep(delay)
|
||||
|
||||
// Retry failed tunnels
|
||||
for tunnelID := range tunnelFailed {
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err == nil {
|
||||
delete(tunnelFailed, tunnelID)
|
||||
fmt.Printf("post-upgrade redeploy retry: tunnel %d succeeded on node %d (attempt %d)\n", tunnelID, nodeID, attempt)
|
||||
} else if !isRetryableError(err) {
|
||||
delete(tunnelFailed, tunnelID) // Non-retryable, don't retry again
|
||||
} else {
|
||||
fmt.Printf("post-upgrade redeploy retry: tunnel %d still failing on node %d (attempt %d): %v\n", tunnelID, nodeID, attempt, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Retry failed forwards
|
||||
var stillFailed []failedForward
|
||||
for _, ff := range failedForwards {
|
||||
if _, skipped := tunnelFailed[ff.forward.TunnelID]; skipped {
|
||||
stillFailed = append(stillFailed, ff) // Tunnel still failed, skip forward
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(ff.forward, "UpdateService", true); err == nil {
|
||||
fmt.Printf("post-upgrade redeploy retry: forward %d succeeded on node %d (attempt %d)\n", ff.id, nodeID, attempt)
|
||||
} else if !isRetryableError(err) {
|
||||
// Non-retryable, drop it
|
||||
} else {
|
||||
stillFailed = append(stillFailed, ff)
|
||||
fmt.Printf("post-upgrade redeploy retry: forward %d still failing on node %d (attempt %d): %v\n", ff.id, nodeID, attempt, err)
|
||||
}
|
||||
}
|
||||
failedForwards = stillFailed
|
||||
|
||||
if len(tunnelFailed) == 0 && len(failedForwards) == 0 {
|
||||
fmt.Printf("post-upgrade redeploy retry: all items recovered on node %d\n", nodeID)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Final summary
|
||||
for tunnelID := range tunnelFailed {
|
||||
fmt.Printf("post-upgrade redeploy: tunnel %d permanently failed on node %d after retries\n", tunnelID, nodeID)
|
||||
}
|
||||
for _, ff := range failedForwards {
|
||||
fmt.Printf("post-upgrade redeploy: forward %d permanently failed on node %d after retries\n", ff.id, nodeID)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -30,21 +30,22 @@ func (User) TableName() string { return "user" }
|
||||
|
||||
// Forward maps to the "forward" table.
|
||||
type Forward struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
UserID int64 `gorm:"column:user_id;not null"`
|
||||
UserName string `gorm:"column:user_name;type:varchar(100);not null"`
|
||||
Name string `gorm:"type:varchar(100);not null"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
|
||||
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
|
||||
InFlow int64 `gorm:"not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
Status int `gorm:"not null"`
|
||||
Inx int `gorm:"not null;default:0"`
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
MaxConn int `gorm:"column:max_conn;not null;default:0"`
|
||||
ProxyProtocol int `gorm:"column:proxy_protocol;not null;default:0"`
|
||||
}
|
||||
|
||||
func (Forward) TableName() string { return "forward" }
|
||||
@@ -439,9 +440,10 @@ type ForwardBackup struct {
|
||||
CreatedTime int64 `json:"createdTime"`
|
||||
UpdatedTime int64 `json:"updatedTime"`
|
||||
Status int `json:"status"`
|
||||
Inx int `json:"inx"`
|
||||
SpeedID *int64 `json:"speedId,omitempty"`
|
||||
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||
Inx int `json:"inx"`
|
||||
SpeedID *int64 `json:"speedId,omitempty"`
|
||||
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
|
||||
ProxyProtocol int `json:"proxyProtocol"`
|
||||
}
|
||||
|
||||
type ForwardPortBackup struct {
|
||||
@@ -537,9 +539,10 @@ type ForwardRecord struct {
|
||||
TunnelID int64
|
||||
RemoteAddr string
|
||||
Strategy string
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
Status int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
ProxyProtocol int
|
||||
}
|
||||
|
||||
// TunnelRecord is a minimal tunnel view used by control plane.
|
||||
|
||||
@@ -295,6 +295,17 @@ func prepareSQLiteLegacyColumns(db *gorm.DB) error {
|
||||
}
|
||||
}
|
||||
|
||||
if m.HasTable(&model.Forward{}) {
|
||||
for _, field := range []string{"ProxyProtocol"} {
|
||||
if m.HasColumn(&model.Forward{}, field) {
|
||||
continue
|
||||
}
|
||||
if err := m.AddColumn(&model.Forward{}, field); err != nil {
|
||||
return fmt.Errorf("add forward.%s: %w", field, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -714,6 +725,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
||||
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
|
||||
"updatedTime": nullableInt64(u.UpdatedTime),
|
||||
"inFlow": u.InFlow, "outFlow": u.OutFlow,
|
||||
"maxConn": u.MaxConn,
|
||||
}
|
||||
if quota := quotaMap[u.ID]; quota != nil {
|
||||
item["dailyQuotaGB"] = quota.DailyLimitGB
|
||||
@@ -769,11 +781,13 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
Status int
|
||||
Inx int
|
||||
SpeedID sql.NullInt64
|
||||
MaxConn int
|
||||
ProxyProtocol int
|
||||
}
|
||||
|
||||
var rows []fwdRow
|
||||
err := r.db.Model(&model.Forward{}).
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id").
|
||||
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.proxy_protocol").
|
||||
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
|
||||
Order("forward.inx ASC, forward.id ASC").
|
||||
Find(&rows).Error
|
||||
@@ -795,6 +809,8 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
|
||||
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
|
||||
"inFlow": row.InFlow, "outFlow": row.OutFlow,
|
||||
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
|
||||
"maxConn": row.MaxConn,
|
||||
"proxyProtocol": row.ProxyProtocol,
|
||||
}
|
||||
if row.SpeedID.Valid {
|
||||
item["speedId"] = row.SpeedID.Int64
|
||||
@@ -1987,6 +2003,7 @@ func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
|
||||
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
|
||||
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
|
||||
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
}
|
||||
ports, err := r.exportForwardPorts(f.ID)
|
||||
if err != nil {
|
||||
@@ -2384,12 +2401,13 @@ func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int
|
||||
UpdatedTime: now,
|
||||
Status: f.Status,
|
||||
Inx: f.Inx,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
}
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy",
|
||||
"in_flow", "out_flow", "updated_time", "status", "inx",
|
||||
"in_flow", "out_flow", "updated_time", "status", "inx", "proxy_protocol",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
if err != nil {
|
||||
|
||||
@@ -124,16 +124,17 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
|
||||
return nil, err
|
||||
}
|
||||
fr := model.ForwardRecord{
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
ID: f.ID,
|
||||
UserID: f.UserID,
|
||||
UserName: f.UserName,
|
||||
Name: f.Name,
|
||||
TunnelID: f.TunnelID,
|
||||
RemoteAddr: f.RemoteAddr,
|
||||
Strategy: f.Strategy,
|
||||
Status: f.Status,
|
||||
SpeedID: f.SpeedID,
|
||||
MaxConn: f.MaxConn,
|
||||
ProxyProtocol: f.ProxyProtocol,
|
||||
}
|
||||
if strings.TrimSpace(fr.Strategy) == "" {
|
||||
fr.Strategy = "fifo"
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestGetForwardRecordIncludesProxyProtocol(t *testing.T) {
|
||||
r, err := Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := r.DB().Create(&model.Forward{
|
||||
UserID: 1,
|
||||
UserName: "admin",
|
||||
Name: "proxy-forward",
|
||||
TunnelID: 1,
|
||||
RemoteAddr: "1.1.1.1:443",
|
||||
Strategy: "fifo",
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
ProxyProtocol: 2,
|
||||
}).Error; err != nil {
|
||||
t.Fatalf("create forward: %v", err)
|
||||
}
|
||||
|
||||
forwardID := mustRepoLastInsertID(t, r)
|
||||
record, err := r.GetForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetForwardRecord: %v", err)
|
||||
}
|
||||
if record == nil {
|
||||
t.Fatalf("expected forward record")
|
||||
}
|
||||
if record.ProxyProtocol != 2 {
|
||||
t.Fatalf("expected proxyProtocol 2, got %d", record.ProxyProtocol)
|
||||
}
|
||||
if record.MaxConn != 0 {
|
||||
t.Fatalf("expected default maxConn 0, got %d", record.MaxConn)
|
||||
}
|
||||
}
|
||||
|
||||
func mustRepoLastInsertID(t *testing.T, r *Repository) int64 {
|
||||
t.Helper()
|
||||
var id int64
|
||||
if err := r.DB().Raw("SELECT last_insert_rowid()").Row().Scan(&id); err != nil {
|
||||
t.Fatalf("last_insert_rowid: %v", err)
|
||||
}
|
||||
if id <= 0 {
|
||||
t.Fatalf("invalid last_insert_rowid %d", id)
|
||||
}
|
||||
return id
|
||||
}
|
||||
@@ -695,20 +695,21 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
|
||||
return p
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int) error {
|
||||
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, proxyProtocol int) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
return r.db.Model(&model.Forward{}).
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"updated_time": now,
|
||||
"name": name,
|
||||
"tunnel_id": tunnelID,
|
||||
"remote_addr": remoteAddr,
|
||||
"strategy": strategy,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
@@ -782,7 +783,7 @@ func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int,
|
||||
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, now int64) {
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, proxyProtocol int, now int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
@@ -798,6 +799,7 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
|
||||
"status": status,
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"max_conn": maxConn,
|
||||
"proxy_protocol": proxyProtocol,
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
@@ -1258,27 +1260,28 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
return ut.ID, true, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int) (int64, error) {
|
||||
func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int, proxyProtocol int) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
var forwardID int64
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
fwd := model.Forward{
|
||||
UserID: userID,
|
||||
UserName: userName,
|
||||
Name: name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: remoteAddr,
|
||||
Strategy: strategy,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
UserID: userID,
|
||||
UserName: userName,
|
||||
Name: name,
|
||||
TunnelID: tunnelID,
|
||||
RemoteAddr: remoteAddr,
|
||||
Strategy: strategy,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
CreatedTime: now,
|
||||
UpdatedTime: now,
|
||||
Status: 1,
|
||||
Inx: inx,
|
||||
MaxConn: maxConn,
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
ProxyProtocol: proxyProtocol,
|
||||
}
|
||||
if err := tx.Create(&fwd).Error; err != nil {
|
||||
return err
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/handler"
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
func TestForwardLocalRemoteAddrToggleContracts(t *testing.T) {
|
||||
handler.DisableSafeRemoteAddrCheckForTesting = false
|
||||
t.Cleanup(func() {
|
||||
handler.DisableSafeRemoteAddrCheckForTesting = true
|
||||
})
|
||||
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(2, 'local_remote_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "local-remote-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "local-remote-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "local-remote-entry", "local-remote-secret", "10.60.0.1", "10.60.0.1", "", "31000-31010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
entryNodeID := mustLastInsertID(t, repo, "local-remote-entry")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 31001, 'round', 1, 'tls')
|
||||
`, tunnelID, entryNodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(601, 2, ?, NULL, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
userToken, err := auth.GenerateToken(2, "local_remote_user", 1, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate user token: %v", err)
|
||||
}
|
||||
|
||||
stopNode := startMockNodeSession(t, server.URL, "local-remote-secret")
|
||||
defer stopNode()
|
||||
waitNodeStatus(t, repo, entryNodeID, 1)
|
||||
|
||||
t.Run("local remote address is rejected on create when toggle is off", func(t *testing.T) {
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "deny-local-create",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected local remote address to be rejected when toggle is off")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "internal IP") && !strings.Contains(out.Msg, "内部") {
|
||||
t.Fatalf("expected internal IP error, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("local remote address is allowed on create when toggle is on", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, "allow_local_remote_addr", "true", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("enable allow_local_remote_addr: %v", err)
|
||||
}
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "allow-local-create",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8080",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
req.Header.Set("Authorization", userToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
assertCode(t, res, 0)
|
||||
})
|
||||
|
||||
t.Run("local remote address is rejected on update when toggle is off", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, "allow_local_remote_addr", "false", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("disable allow_local_remote_addr: %v", err)
|
||||
}
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "safe-remote-before-update",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "8.8.8.8:53",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal safe create payload: %v", err)
|
||||
}
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", userToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "safe-remote-before-update")
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "safe-remote-before-update",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8081",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", userToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(updateRes.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode update response: %v", err)
|
||||
}
|
||||
if out.Code == 0 {
|
||||
t.Fatalf("expected local remote address to be rejected on update when toggle is off")
|
||||
}
|
||||
if !strings.Contains(out.Msg, "internal IP") && !strings.Contains(out.Msg, "内部") {
|
||||
t.Fatalf("expected internal IP error on update, got code=%d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("local remote address is allowed on update when toggle is on", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO vite_config(name, value, time)
|
||||
VALUES(?, ?, ?)
|
||||
ON CONFLICT(name) DO UPDATE SET value = excluded.value, time = excluded.time
|
||||
`, "allow_local_remote_addr", "true", time.Now().UnixMilli()).Error; err != nil {
|
||||
t.Fatalf("enable allow_local_remote_addr: %v", err)
|
||||
}
|
||||
|
||||
var forwardID int64
|
||||
if err := repo.DB().Raw(`SELECT id FROM forward WHERE name = ? ORDER BY id DESC LIMIT 1`, "safe-remote-before-update").Row().Scan(&forwardID); err != nil {
|
||||
t.Fatalf("query forward id: %v", err)
|
||||
}
|
||||
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "safe-remote-before-update",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "127.0.0.1:8081",
|
||||
"strategy": "fifo",
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", userToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
assertCode(t, updateRes, 0)
|
||||
})
|
||||
}
|
||||
@@ -96,6 +96,7 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"maxConn": 42,
|
||||
"proxyProtocol": 2,
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
@@ -120,6 +121,47 @@ func TestMaxConnLimit(t *testing.T) {
|
||||
t.Fatalf("get forward ID: %v", err)
|
||||
}
|
||||
|
||||
listOut := requestContractEnvelope(t, router, adminToken, "/api/v1/forward/list", nil)
|
||||
if listOut.Code != 0 {
|
||||
t.Fatalf("expected /forward/list success, got code=%d msg=%s", listOut.Code, listOut.Msg)
|
||||
}
|
||||
|
||||
rows := mustContractSlice(t, listOut.Data, "forward list")
|
||||
var target map[string]interface{}
|
||||
for _, row := range rows {
|
||||
item, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected forward item to be object, got %T", row)
|
||||
}
|
||||
idVal, ok := item["id"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected forward id to be float64, got %T", item["id"])
|
||||
}
|
||||
if int64(idVal) == forwardID {
|
||||
target = item
|
||||
break
|
||||
}
|
||||
}
|
||||
if target == nil {
|
||||
t.Fatalf("forward %d not found in /forward/list response", forwardID)
|
||||
}
|
||||
|
||||
maxConnVal, ok := target["maxConn"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected maxConn to be float64, got %T (%v)", target["maxConn"], target["maxConn"])
|
||||
}
|
||||
if int(maxConnVal) != 42 {
|
||||
t.Fatalf("expected maxConn 42 in /forward/list, got %v", maxConnVal)
|
||||
}
|
||||
|
||||
proxyProtocolVal, ok := target["proxyProtocol"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected proxyProtocol to be float64, got %T (%v)", target["proxyProtocol"], target["proxyProtocol"])
|
||||
}
|
||||
if int(proxyProtocolVal) != 2 {
|
||||
t.Fatalf("expected proxyProtocol 2 in /forward/list, got %v", proxyProtocolVal)
|
||||
}
|
||||
|
||||
commandMu.Lock()
|
||||
defer commandMu.Unlock()
|
||||
|
||||
|
||||
@@ -349,9 +349,9 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
tunnelID := mustLastInsertID(t, r, "backup-forward-tunnel")
|
||||
|
||||
if err := r.DB().Exec(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88).Error; err != nil {
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx, proxy_protocol)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88, 2).Error; err != nil {
|
||||
t.Fatalf("seed forward for backup: %v", err)
|
||||
}
|
||||
forwardID := mustLastInsertID(t, r, "backup-forward")
|
||||
@@ -412,6 +412,9 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
if !ok {
|
||||
t.Fatalf("expected forwardPorts for forward %d in payload", forwardID)
|
||||
}
|
||||
if proxyProtocol, ok := forwardMap["proxyProtocol"].(float64); !ok || int(proxyProtocol) != 2 {
|
||||
t.Fatalf("expected exported proxyProtocol 2 for forward %d, got %v", forwardID, forwardMap["proxyProtocol"])
|
||||
}
|
||||
for _, p := range portsRaw {
|
||||
portMap, ok := p.(map[string]interface{})
|
||||
if !ok {
|
||||
@@ -475,6 +478,14 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
|
||||
t.Fatalf("expected forward_port node=%d port=%d after import, got %v", nodeID, port, after)
|
||||
}
|
||||
}
|
||||
|
||||
var proxyProtocol int
|
||||
if err := r.DB().Raw(`SELECT proxy_protocol FROM forward WHERE id = ?`, forwardID).Row().Scan(&proxyProtocol); err != nil {
|
||||
t.Fatalf("query proxy_protocol after import: %v", err)
|
||||
}
|
||||
if proxyProtocol != 2 {
|
||||
t.Fatalf("expected proxy_protocol 2 after import, got %d", proxyProtocol)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package contract_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
)
|
||||
|
||||
func TestUserListReturnsMaxConn(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, max_conn, created_time, updated_time, status)
|
||||
VALUES(2, 'max_conn_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 10, 37, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
out := requestContractEnvelope(t, router, adminToken, "/api/v1/user/list", map[string]interface{}{})
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected /user/list success, got code=%d msg=%s", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
rows := mustContractSlice(t, out.Data, "user list")
|
||||
var target map[string]interface{}
|
||||
for _, row := range rows {
|
||||
item, ok := row.(map[string]interface{})
|
||||
if !ok {
|
||||
t.Fatalf("expected user item to be object, got %T", row)
|
||||
}
|
||||
idVal, ok := item["id"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected user id to be float64, got %T", item["id"])
|
||||
}
|
||||
if int64(idVal) == 2 {
|
||||
target = item
|
||||
break
|
||||
}
|
||||
}
|
||||
if target == nil {
|
||||
t.Fatalf("user 2 not found in /user/list response")
|
||||
}
|
||||
|
||||
maxConnVal, ok := target["maxConn"].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("expected maxConn to be float64, got %T (%v)", target["maxConn"], target["maxConn"])
|
||||
}
|
||||
if int(maxConnVal) != 37 {
|
||||
t.Fatalf("expected maxConn 37 in /user/list, got %v", maxConnVal)
|
||||
}
|
||||
}
|
||||
@@ -116,6 +116,10 @@ func main() {
|
||||
|
||||
fmt.Printf("✅ 配置加载成功 - addr: %s\n", config.Addr)
|
||||
|
||||
// 设置运行时配置持久化路径
|
||||
socket.SetConfigPersistPath("gost.json")
|
||||
// 启用持久化将在 program.Start() 后开启,避免启动加载阶段触发冗余写入
|
||||
|
||||
log := xlogger.NewLogger()
|
||||
logger.SetDefault(log)
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
metrics "github.com/go-gost/x/metrics/service"
|
||||
"github.com/go-gost/x/registry"
|
||||
xservice "github.com/go-gost/x/service"
|
||||
"github.com/go-gost/x/socket"
|
||||
"github.com/judwhite/go-svc"
|
||||
"net/http"
|
||||
"os"
|
||||
@@ -66,6 +67,10 @@ func (p *program) Start() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// Enable config persistence after initial load so runtime mutations
|
||||
// (AddService, UpdateService, DeleteService, etc.) are saved to disk.
|
||||
socket.EnableConfigPersist()
|
||||
|
||||
if err := p.run(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -44,9 +44,14 @@ func Set(c *Config) {
|
||||
|
||||
func OnUpdate(f func(c *Config) error) error {
|
||||
globalMux.Lock()
|
||||
defer globalMux.Unlock()
|
||||
err := f(global)
|
||||
globalMux.Unlock()
|
||||
|
||||
return f(global)
|
||||
if err == nil {
|
||||
persist()
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
type LogConfig struct {
|
||||
@@ -573,6 +578,7 @@ func (c *Config) Load() error {
|
||||
if err := v.ReadInConfig(); err != nil {
|
||||
return err
|
||||
}
|
||||
SetPersistPath(v.ConfigFileUsed())
|
||||
|
||||
return v.Unmarshal(c)
|
||||
}
|
||||
@@ -590,6 +596,7 @@ func (c *Config) ReadFile(file string) error {
|
||||
if err := v.ReadInConfig(); err != nil {
|
||||
return err
|
||||
}
|
||||
SetPersistPath(v.ConfigFileUsed())
|
||||
return v.Unmarshal(c)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var (
|
||||
persistPath string
|
||||
persistMu sync.Mutex
|
||||
persistEnable bool
|
||||
)
|
||||
|
||||
// SetPersistPath sets the file path where runtime config changes will be
|
||||
// automatically persisted. Call this once during agent startup before any
|
||||
// OnUpdate mutations occur.
|
||||
func SetPersistPath(path string) {
|
||||
persistMu.Lock()
|
||||
defer persistMu.Unlock()
|
||||
persistPath = path
|
||||
}
|
||||
|
||||
func PersistPath() string {
|
||||
persistMu.Lock()
|
||||
defer persistMu.Unlock()
|
||||
return persistPath
|
||||
}
|
||||
|
||||
// EnablePersist turns on automatic persistence. Call this after the initial
|
||||
// config has been loaded (e.g. after program.Start) so that startup loading
|
||||
// does not trigger redundant disk writes.
|
||||
func EnablePersist() {
|
||||
persistMu.Lock()
|
||||
defer persistMu.Unlock()
|
||||
persistEnable = true
|
||||
}
|
||||
|
||||
// persist writes the current global config to the configured file atomically.
|
||||
func persist() {
|
||||
persistMu.Lock()
|
||||
path := persistPath
|
||||
enabled := persistEnable
|
||||
persistMu.Unlock()
|
||||
|
||||
if !enabled || path == "" {
|
||||
return
|
||||
}
|
||||
|
||||
cfg := Global()
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
enc.SetIndent("", " ")
|
||||
if err := enc.Encode(cfg); err != nil {
|
||||
fmt.Printf("⚠️ config persist: marshal failed: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
// Atomic write: write to temp file then rename
|
||||
dir := filepath.Dir(path)
|
||||
tmp, err := os.CreateTemp(dir, ".gost-*.tmp")
|
||||
if err != nil {
|
||||
fmt.Printf("⚠️ config persist: create temp file failed: %v\n", err)
|
||||
return
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
|
||||
if _, err := tmp.Write(buf.Bytes()); err != nil {
|
||||
tmp.Close()
|
||||
os.Remove(tmpName)
|
||||
fmt.Printf("⚠️ config persist: write failed: %v\n", err)
|
||||
return
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
os.Remove(tmpName)
|
||||
fmt.Printf("⚠️ config persist: close temp file failed: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
os.Remove(tmpName)
|
||||
fmt.Printf("⚠️ config persist: rename failed: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Printf("💾 节点配置已持久化到 %s\n", path)
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestReadFileSetsPersistPath(t *testing.T) {
|
||||
originalPath := persistPath
|
||||
originalEnabled := persistEnable
|
||||
persistPath = ""
|
||||
persistEnable = false
|
||||
t.Cleanup(func() {
|
||||
persistPath = originalPath
|
||||
persistEnable = originalEnabled
|
||||
})
|
||||
|
||||
dir := t.TempDir()
|
||||
configFile := filepath.Join(dir, "custom-gost.yaml")
|
||||
if err := os.WriteFile(configFile, []byte("services: []\n"), 0o644); err != nil {
|
||||
t.Fatalf("write config file: %v", err)
|
||||
}
|
||||
|
||||
var cfg Config
|
||||
if err := cfg.ReadFile(configFile); err != nil {
|
||||
t.Fatalf("ReadFile: %v", err)
|
||||
}
|
||||
if persistPath != configFile {
|
||||
t.Fatalf("expected persistPath %q, got %q", configFile, persistPath)
|
||||
}
|
||||
}
|
||||
@@ -1690,6 +1690,26 @@ func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls
|
||||
return reporter
|
||||
}
|
||||
|
||||
var configPersistPath string
|
||||
|
||||
// SetConfigPersistPath sets the path where runtime config changes will be
|
||||
// persisted to disk (gost.json). Called by main during agent startup.
|
||||
func SetConfigPersistPath(path string) {
|
||||
configPersistPath = path
|
||||
config.SetPersistPath(path)
|
||||
}
|
||||
|
||||
// EnableConfigPersist turns on automatic disk persistence after the initial
|
||||
// config has been loaded and applied.
|
||||
func EnableConfigPersist() {
|
||||
config.EnablePersist()
|
||||
path := config.PersistPath()
|
||||
if path == "" {
|
||||
path = configPersistPath
|
||||
}
|
||||
fmt.Printf("🔒 节点配置持久化已启用,运行时变更将自动保存到 %s\n", path)
|
||||
}
|
||||
|
||||
// handleTcpPing 处理TCP ping诊断命令
|
||||
func (w *WebSocketReporter) handleTcpPing(data interface{}) (TcpPingResponse, error) {
|
||||
jsonData, err := json.Marshal(data)
|
||||
|
||||
+18
-27
@@ -5,7 +5,7 @@ import {
|
||||
useNavigate,
|
||||
Navigate,
|
||||
} from "react-router-dom";
|
||||
import { useEffect, useState } from "react";
|
||||
import { useEffect } from "react";
|
||||
import { AnimatePresence } from "framer-motion";
|
||||
|
||||
import IndexPage from "@/pages/index";
|
||||
@@ -39,32 +39,7 @@ const ProtectedRoute = ({
|
||||
skipLayout?: boolean;
|
||||
}) => {
|
||||
const isH5 = useH5Mode();
|
||||
const navigate = useNavigate();
|
||||
const [authenticated, setAuthenticated] = useState(() => isLoggedIn());
|
||||
|
||||
useEffect(() => {
|
||||
const handleSessionChange = () => {
|
||||
const loggedIn = isLoggedIn();
|
||||
|
||||
setAuthenticated(loggedIn);
|
||||
|
||||
if (!loggedIn) {
|
||||
navigate("/", { replace: true });
|
||||
}
|
||||
};
|
||||
|
||||
window.addEventListener(SESSION_UPDATED_EVENT, handleSessionChange);
|
||||
|
||||
return () => {
|
||||
window.removeEventListener(SESSION_UPDATED_EVENT, handleSessionChange);
|
||||
};
|
||||
}, [navigate]);
|
||||
|
||||
useEffect(() => {
|
||||
if (!authenticated) {
|
||||
navigate("/", { replace: true });
|
||||
}
|
||||
}, [authenticated, navigate]);
|
||||
const authenticated = isLoggedIn();
|
||||
|
||||
if (!authenticated) {
|
||||
return <Navigate replace to="/" />;
|
||||
@@ -103,6 +78,22 @@ const LoginRoute = () => {
|
||||
|
||||
function App() {
|
||||
const location = useLocation();
|
||||
const navigate = useNavigate();
|
||||
|
||||
// 全局登录状态监听,当检测到未登录且不在首页时,跳转到首页
|
||||
useEffect(() => {
|
||||
const handleSessionUpdate = () => {
|
||||
if (!isLoggedIn() && location.pathname !== "/") {
|
||||
navigate("/", { replace: true });
|
||||
}
|
||||
};
|
||||
|
||||
window.addEventListener(SESSION_UPDATED_EVENT, handleSessionUpdate);
|
||||
|
||||
return () => {
|
||||
window.removeEventListener(SESSION_UPDATED_EVENT, handleSessionUpdate);
|
||||
};
|
||||
}, [location.pathname, navigate]);
|
||||
|
||||
// 处理自定义背景图片
|
||||
useEffect(() => {
|
||||
|
||||
@@ -71,6 +71,8 @@ export interface ForwardApiItem {
|
||||
userId?: number;
|
||||
tunnelId?: number;
|
||||
speedId?: number | null;
|
||||
maxConn?: number;
|
||||
proxyProtocol?: number;
|
||||
inx?: number;
|
||||
[key: string]: unknown;
|
||||
}
|
||||
@@ -201,7 +203,7 @@ export interface UserPackageInfoApiData {
|
||||
num: number;
|
||||
expTime?: string;
|
||||
flowResetTime?: number;
|
||||
maxConn?: number;
|
||||
maxConn?: number;
|
||||
[key: string]: unknown;
|
||||
};
|
||||
tunnelPermissions: UserTunnelPermissionApiItem[];
|
||||
@@ -370,6 +372,8 @@ export interface ForwardMutationPayload {
|
||||
remoteAddr?: string;
|
||||
strategy?: string;
|
||||
speedId?: number | null;
|
||||
maxConn?: number;
|
||||
proxyProtocol?: number;
|
||||
}
|
||||
|
||||
export interface SpeedLimitMutationPayload {
|
||||
|
||||
@@ -241,7 +241,6 @@ export default function AdminLayout({
|
||||
// 退出登录
|
||||
const handleLogout = () => {
|
||||
safeLogout();
|
||||
navigate("/");
|
||||
};
|
||||
|
||||
// 切换移动端菜单
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import { useState } from "react";
|
||||
import { useNavigate } from "react-router-dom";
|
||||
import toast from "react-hot-toast";
|
||||
|
||||
import { Button } from "@/shadcn-bridge/heroui/button";
|
||||
@@ -26,7 +25,6 @@ export default function ChangePasswordPage() {
|
||||
});
|
||||
const [loading, setLoading] = useState(false);
|
||||
const [errors, setErrors] = useState<Partial<PasswordForm>>({});
|
||||
const navigate = useNavigate();
|
||||
|
||||
const validateForm = (): boolean => {
|
||||
const newErrors: Partial<PasswordForm> = {};
|
||||
@@ -98,7 +96,6 @@ export default function ChangePasswordPage() {
|
||||
|
||||
const logout = () => {
|
||||
safeLogout();
|
||||
navigate("/");
|
||||
};
|
||||
|
||||
const handleKeyPress = (e: React.KeyboardEvent) => {
|
||||
|
||||
@@ -185,6 +185,13 @@ const CONFIG_ITEMS: ConfigItem[] = [
|
||||
dependsOn: "github_proxy_enabled",
|
||||
dependsValue: "true",
|
||||
},
|
||||
{
|
||||
key: "allow_local_remote_addr",
|
||||
label: "允许转发到本地地址",
|
||||
description:
|
||||
"开启后,普通用户创建或编辑规则时可将目标地址指向 127.0.0.1、10.x.x.x、172.16-31.x.x、192.168.x.x 等本地或内网地址。默认关闭以降低开放代理风险。",
|
||||
type: "switch",
|
||||
},
|
||||
];
|
||||
|
||||
const BACKUP_TYPE_OPTIONS = [
|
||||
@@ -219,6 +226,7 @@ const getInitialConfigs = (): Record<string, string> => {
|
||||
"app_favicon",
|
||||
"github_proxy_enabled",
|
||||
"github_proxy_url",
|
||||
"allow_local_remote_addr",
|
||||
];
|
||||
const initialConfigs: Record<string, string> = {};
|
||||
|
||||
|
||||
@@ -50,7 +50,6 @@ export default function DashboardPage() {
|
||||
const handleLogout = () => {
|
||||
safeLogout();
|
||||
toast.success("已退出登录");
|
||||
navigate("/");
|
||||
};
|
||||
|
||||
const {
|
||||
|
||||
@@ -54,6 +54,7 @@ import { Switch } from "@/shadcn-bridge/heroui/switch";
|
||||
import { Alert } from "@/shadcn-bridge/heroui/alert";
|
||||
import { Progress } from "@/shadcn-bridge/heroui/progress";
|
||||
import { Checkbox } from "@/shadcn-bridge/heroui/checkbox";
|
||||
import { Accordion, AccordionItem } from "@/shadcn-bridge/heroui/accordion";
|
||||
import {
|
||||
createForward,
|
||||
getForwardList,
|
||||
@@ -118,11 +119,13 @@ interface Forward {
|
||||
outFlow: number;
|
||||
serviceRunning: boolean;
|
||||
federationShareFlow?: number;
|
||||
maxConn?: number;
|
||||
createdTime: string;
|
||||
userName?: string;
|
||||
userId?: number;
|
||||
inx?: number;
|
||||
speedId?: number | null;
|
||||
proxyProtocol?: number;
|
||||
}
|
||||
|
||||
interface Tunnel {
|
||||
@@ -158,6 +161,7 @@ interface ForwardForm {
|
||||
strategy: string;
|
||||
speedId: number | null;
|
||||
maxConn?: number;
|
||||
proxyProtocol?: number;
|
||||
}
|
||||
|
||||
interface ForwardUserGroup {
|
||||
@@ -574,6 +578,11 @@ const mapForwardApiItems = (items: ForwardApiItem[]): Forward[] => {
|
||||
typeof forward.speedId === "number" || forward.speedId === null
|
||||
? forward.speedId
|
||||
: undefined,
|
||||
maxConn: typeof forward.maxConn === "number" ? forward.maxConn : undefined,
|
||||
proxyProtocol:
|
||||
typeof forward.proxyProtocol === "number"
|
||||
? forward.proxyProtocol
|
||||
: undefined,
|
||||
serviceRunning: forward.status === 1,
|
||||
}));
|
||||
};
|
||||
@@ -1308,6 +1317,7 @@ export default function ForwardPage() {
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
maxConn: 0,
|
||||
proxyProtocol: 0,
|
||||
});
|
||||
const [inIpTouched, setInIpTouched] = useState(false);
|
||||
|
||||
@@ -2095,6 +2105,7 @@ export default function ForwardPage() {
|
||||
interfaceName: "",
|
||||
strategy: "fifo",
|
||||
speedId: null,
|
||||
proxyProtocol: 0,
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -2115,6 +2126,8 @@ export default function ForwardPage() {
|
||||
interfaceName: forward.interfaceName || "",
|
||||
strategy: forward.strategy || "fifo",
|
||||
speedId: normalizeSpeedId(forward.speedId),
|
||||
maxConn: forward.maxConn ?? 0,
|
||||
proxyProtocol: forward.proxyProtocol ?? 0,
|
||||
});
|
||||
setErrors({});
|
||||
setModalOpen(true);
|
||||
@@ -2244,6 +2257,7 @@ export default function ForwardPage() {
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizedSpeedId,
|
||||
maxConn: form.maxConn,
|
||||
proxyProtocol: form.proxyProtocol,
|
||||
};
|
||||
|
||||
res = await updateForward(updateData);
|
||||
@@ -2257,11 +2271,11 @@ export default function ForwardPage() {
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizedSpeedId,
|
||||
maxConn: form.maxConn,
|
||||
proxyProtocol: form.proxyProtocol,
|
||||
};
|
||||
|
||||
res = await createForward(createData);
|
||||
}
|
||||
|
||||
if (res.code === 0) {
|
||||
const warningItems = Array.isArray((res as any).data?.warnings)
|
||||
? (res as any).data.warnings
|
||||
@@ -4727,20 +4741,6 @@ export default function ForwardPage() {
|
||||
</ModalHeader>
|
||||
<ModalBody>
|
||||
<div className="space-y-4 pb-4">
|
||||
<Input
|
||||
label="最大连接数"
|
||||
placeholder="0 或空表示不限制"
|
||||
type="number"
|
||||
min="0"
|
||||
value={form.maxConn === 0 ? "" : String(form.maxConn || "")}
|
||||
onChange={(e) => {
|
||||
const value = Math.max(Number(e.target.value) || 0, 0);
|
||||
setForm((prev) => ({ ...prev, maxConn: value }));
|
||||
}}
|
||||
description="此设置优先于用户的全局连接数限制。0 表示不限制。"
|
||||
variant="bordered"
|
||||
/>
|
||||
|
||||
<Input
|
||||
errorMessage={errors.name}
|
||||
isInvalid={!!errors.name}
|
||||
@@ -4753,38 +4753,6 @@ export default function ForwardPage() {
|
||||
}
|
||||
/>
|
||||
|
||||
{isAdmin && (
|
||||
<Select
|
||||
label="规则限速"
|
||||
placeholder="不限速"
|
||||
selectedKeys={
|
||||
selectedSpeedId !== null
|
||||
? [selectedSpeedId.toString()]
|
||||
: []
|
||||
}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as
|
||||
| string
|
||||
| undefined;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
speedId: selectedKey ? Number(selectedKey) : null,
|
||||
}));
|
||||
}}
|
||||
>
|
||||
{availableSpeedLimits.map((speedLimit) => (
|
||||
<SelectItem
|
||||
key={speedLimit.id.toString()}
|
||||
textValue={speedLimit.name}
|
||||
>
|
||||
{speedLimit.name}
|
||||
</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
)}
|
||||
|
||||
<Select
|
||||
description={
|
||||
isEdit
|
||||
@@ -4912,6 +4880,91 @@ export default function ForwardPage() {
|
||||
<SelectItem key="hash">哈希模式 - IP哈希</SelectItem>
|
||||
</Select>
|
||||
)}
|
||||
<Accordion className="px-0" variant="light">
|
||||
<AccordionItem
|
||||
key="advanced"
|
||||
aria-label="高级设置"
|
||||
title={
|
||||
<span className="text-small text-default-500 font-medium">
|
||||
高级设置
|
||||
</span>
|
||||
}
|
||||
>
|
||||
<div className="space-y-4 pb-2">
|
||||
<Input
|
||||
description="此设置优先于用户的全局连接数限制。0 表示不限制。"
|
||||
label="最大连接数"
|
||||
min="0"
|
||||
placeholder="0 或空表示不限制"
|
||||
type="number"
|
||||
value={
|
||||
form.maxConn === 0 ? "" : String(form.maxConn || "")
|
||||
}
|
||||
variant="bordered"
|
||||
onChange={(e) => {
|
||||
const value = Math.max(
|
||||
Number(e.target.value) || 0,
|
||||
0,
|
||||
);
|
||||
|
||||
setForm((prev) => ({ ...prev, maxConn: value }));
|
||||
}}
|
||||
/>
|
||||
<Select
|
||||
description="启用 PROXY protocol,用于透传客户端真实 IP"
|
||||
label="Proxy Protocol"
|
||||
placeholder="禁用"
|
||||
selectedKeys={[String(form.proxyProtocol || 0)]}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as string;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
proxyProtocol: Number(selectedKey),
|
||||
}));
|
||||
}}
|
||||
>
|
||||
<SelectItem key="0">禁用</SelectItem>
|
||||
<SelectItem key="1">Version 1</SelectItem>
|
||||
<SelectItem key="2">Version 2</SelectItem>
|
||||
</Select>
|
||||
{isAdmin && (
|
||||
<Select
|
||||
label="规则限速"
|
||||
placeholder="不限速"
|
||||
selectedKeys={
|
||||
selectedSpeedId !== null
|
||||
? [selectedSpeedId.toString()]
|
||||
: []
|
||||
}
|
||||
variant="bordered"
|
||||
onSelectionChange={(keys) => {
|
||||
const selectedKey = Array.from(keys)[0] as
|
||||
| string
|
||||
| undefined;
|
||||
|
||||
setForm((prev) => ({
|
||||
...prev,
|
||||
speedId: selectedKey
|
||||
? Number(selectedKey)
|
||||
: null,
|
||||
}));
|
||||
}}
|
||||
>
|
||||
{availableSpeedLimits.map((speedLimit) => (
|
||||
<SelectItem
|
||||
key={speedLimit.id.toString()}
|
||||
textValue={speedLimit.name}
|
||||
>
|
||||
{speedLimit.name}
|
||||
</SelectItem>
|
||||
))}
|
||||
</Select>
|
||||
)}
|
||||
</div>
|
||||
</AccordionItem>
|
||||
</Accordion>
|
||||
</div>
|
||||
</ModalBody>
|
||||
<ModalFooter>
|
||||
|
||||
@@ -127,7 +127,6 @@ export default function ProfilePage() {
|
||||
// 退出登录
|
||||
const handleLogout = () => {
|
||||
safeLogout();
|
||||
navigate("/", { replace: true });
|
||||
};
|
||||
|
||||
// 密码表单验证
|
||||
|
||||
@@ -160,6 +160,7 @@ const normalizeUserItem = (item: Partial<User>): User => {
|
||||
monthlyUsedBytes: Number(item.monthlyUsedBytes ?? 0),
|
||||
disabledByQuota: Number(item.disabledByQuota ?? 0),
|
||||
quotaDisabledAt: Number(item.quotaDisabledAt ?? 0),
|
||||
maxConn: item.maxConn != null ? Number(item.maxConn) : undefined,
|
||||
};
|
||||
};
|
||||
|
||||
@@ -515,7 +516,7 @@ export default function UserPage() {
|
||||
num: 10,
|
||||
expTime: null,
|
||||
flowResetTime: 0,
|
||||
maxConn: 0,
|
||||
maxConn: 0,
|
||||
groupIds: [],
|
||||
});
|
||||
onUserModalOpen();
|
||||
@@ -545,6 +546,7 @@ export default function UserPage() {
|
||||
num: user.num,
|
||||
expTime: user.expTime ? new Date(user.expTime) : null,
|
||||
flowResetTime: user.flowResetTime ?? 0,
|
||||
maxConn: user.maxConn ?? 0,
|
||||
groupIds: currentGroupIds,
|
||||
});
|
||||
onUserModalOpen();
|
||||
@@ -1520,23 +1522,15 @@ export default function UserPage() {
|
||||
/>
|
||||
<Input
|
||||
label="最大连接数"
|
||||
min="0"
|
||||
placeholder="0 或空表示不限制"
|
||||
type="number"
|
||||
min="0"
|
||||
value={userForm.maxConn === 0 ? "" : String(userForm.maxConn || "")}
|
||||
onChange={(e) => {
|
||||
const value = Math.max(Number(e.target.value) || 0, 0);
|
||||
setUserForm((prev) => ({ ...prev, maxConn: value }));
|
||||
}}
|
||||
/>
|
||||
<Input
|
||||
label="最大连接数"
|
||||
placeholder="0 或空表示不限制"
|
||||
type="number"
|
||||
min="0"
|
||||
value={userForm.maxConn === 0 ? "" : String(userForm.maxConn || "")}
|
||||
value={
|
||||
userForm.maxConn === 0 ? "" : String(userForm.maxConn || "")
|
||||
}
|
||||
onChange={(e) => {
|
||||
const value = Math.max(Number(e.target.value) || 0, 0);
|
||||
|
||||
setUserForm((prev) => ({ ...prev, maxConn: value }));
|
||||
}}
|
||||
/>
|
||||
|
||||
@@ -24,6 +24,7 @@ export interface User {
|
||||
monthlyUsedBytes?: number;
|
||||
disabledByQuota?: number;
|
||||
quotaDisabledAt?: number;
|
||||
maxConn?: number;
|
||||
}
|
||||
|
||||
export interface UserGroup {
|
||||
|
||||
@@ -2,7 +2,7 @@ import { clearSession } from "@/utils/session";
|
||||
|
||||
/**
|
||||
* 安全退出登录函数
|
||||
* 清除登录相关数据,但保留用户偏好设置(如主题)
|
||||
* 清除登录相关数据
|
||||
*/
|
||||
export const safeLogout = () => {
|
||||
clearSession();
|
||||
|
||||
Reference in New Issue
Block a user