Compare commits

...

3 Commits

Author SHA1 Message Date
sagitchu a5c33c68d9 Revert "fix: 保留用户/规则高级设置并增强节点运行时恢复 (#468)"
This reverts commit 3f374df724.
2026-04-25 22:45:27 +08:00
sagitchu 325baa8adc Revert "feat: 批量设置用户最大连接数"
This reverts commit 28568244bb.
2026-04-25 22:45:23 +08:00
sagitchu 28568244bb feat: 批量设置用户最大连接数
前端:
- 用户列表添加 checkbox 多选(表格和网格视图)
- 选中用户后显示「批量设置连接数」按钮
- 批量设置对话框,输入连接数后批量应用

后端:
- 新增 BatchUpdateUserMaxConn repo 方法
- 新增 /api/v1/user/batch-set-max-conn API
- 返回成功/失败计数
- 加入 admin-only 路由白名单
2026-04-25 22:03:42 +08:00
30 changed files with 123 additions and 744 deletions
+1 -1
View File
@@ -90,7 +90,7 @@
| 表名 | Model | 特殊处理 | | 表名 | Model | 特殊处理 |
|------|-------|----------| |------|-------|----------|
| `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) | | `user` | `User` | `TableName()` 返回 `"user"` (PG 保留字) |
| `forward` | `Forward` | 增加 `proxy_protocol` 字段 | | `forward` | `Forward` | |
| `forward_port` | `ForwardPort` | | | `forward_port` | `ForwardPort` | |
| `node` | `Node` | | | `node` | `Node` | |
| `speed_limit` | `SpeedLimit` | | | `speed_limit` | `SpeedLimit` | |
@@ -1705,12 +1705,6 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
if cLimiterName != "" { if cLimiterName != "" {
service["climiter"] = 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" { if protocol == "udp" {
listenerMetadata := map[string]interface{}{ listenerMetadata := map[string]interface{}{
"keepAlive": true, "keepAlive": true,
@@ -1722,10 +1716,7 @@ func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel
service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID) service["handler"].(map[string]interface{})["chain"] = fmt.Sprintf("chains_%d", forward.TunnelID)
} }
if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" { if tunnel != nil && tunnel.Type == 1 && strings.TrimSpace(node.InterfaceName) != "" {
if service["metadata"] == nil { service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
service["metadata"] = map[string]interface{}{}
}
service["metadata"].(map[string]interface{})["interface"] = node.InterfaceName
} }
if limiterID != nil && *limiterID > 0 { if limiterID != nil && *limiterID > 0 {
service["limiter"] = strconv.FormatInt(*limiterID, 10) service["limiter"] = strconv.FormatInt(*limiterID, 10)
@@ -1,98 +0,0 @@
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)
}
}
@@ -1728,7 +1728,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
} }
if roleID != 0 { if roleID != 0 {
if err := IsSafeRemoteAddr(remoteAddr); err != nil { if err := IsSafeRemoteAddr(remoteAddr); err != nil {
response.WriteJSON(w, response.Err(403, err.Error())) response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络"))
return return
} }
if speedIDVal, ok := req["speedId"]; ok && speedIDVal != nil { if speedIDVal, ok := req["speedId"]; ok && speedIDVal != nil {
@@ -1780,9 +1780,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
userName = "user" userName = "user"
} }
maxConn := asInt(req["maxConn"], 0) maxConn := asInt(req["maxConn"], 0)
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)
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 { if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
@@ -1855,7 +1853,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
} }
if actorRole != 0 { if actorRole != 0 {
if err := IsSafeRemoteAddr(remoteAddr); err != nil { if err := IsSafeRemoteAddr(remoteAddr); err != nil {
response.WriteJSON(w, response.Err(403, err.Error())) response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络"))
return return
} }
} }
@@ -1934,9 +1932,8 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
} }
now := time.Now().UnixMilli() now := time.Now().UnixMilli()
maxConn := asInt(req["maxConn"], forward.MaxConn) 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, proxyProtocol); err != nil { if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID, maxConn); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
return return
} }
@@ -4088,7 +4085,7 @@ func (h *Handler) rollbackForwardMutation(oldForward *forwardRecord, oldPorts []
h.repo.RollbackForwardFields( h.repo.RollbackForwardFields(
oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name, oldForward.ID, oldForward.UserID, oldForward.UserName, oldForward.Name,
oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status, oldForward.TunnelID, oldForward.RemoteAddr, oldForward.Strategy, oldForward.Status,
oldForward.SpeedID, oldForward.MaxConn, oldForward.ProxyProtocol, oldForward.SpeedID, oldForward.MaxConn,
time.Now().UnixMilli(), time.Now().UnixMilli(),
) )
@@ -11,37 +11,11 @@ var DisableSafeRemoteAddrCheckForTesting = false
// IsSafeRemoteAddr checks if a given address is safe to connect to (prevents SSRF/Open Proxy). // 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. // 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 { func IsSafeRemoteAddr(addr string) error {
if DisableSafeRemoteAddrCheckForTesting { if DisableSafeRemoteAddrCheckForTesting {
return nil 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) host, _, err := net.SplitHostPort(addr)
if err != nil { if err != nil {
if strings.Contains(err.Error(), "missing port in address") { if strings.Contains(err.Error(), "missing port in address") {
@@ -53,12 +27,12 @@ func checkSingleRemoteAddr(addr string) error {
ips, err := net.LookupIP(host) ips, err := net.LookupIP(host)
if err != nil { if err != nil {
return fmt.Errorf("could not resolve address %q: %v", addr, err) return fmt.Errorf("could not resolve address: %v", err)
} }
for _, ip := range ips { for _, ip := range ips {
if ip.IsLoopback() || ip.IsPrivate() { if ip.IsLoopback() || ip.IsPrivate() {
return fmt.Errorf("address %q resolves to internal IP: %s", addr, ip.String()) return fmt.Errorf("address resolves to internal IP: %s", ip.String())
} }
} }
@@ -13,13 +13,6 @@ import (
"go-backend/internal/http/response" "go-backend/internal/http/response"
) )
// failedForward tracks a forward that failed redeployment, for retry.
type failedForward struct {
id int64
forward *forwardRecord
err error
}
const ( const (
githubRepo = "Sagit-chu/flvx" githubRepo = "Sagit-chu/flvx"
githubAPIBase = "https://api.github.com" githubAPIBase = "https://api.github.com"
@@ -397,9 +390,6 @@ func (h *Handler) consumeNodePendingUpgradeRedeploy(nodeID int64) bool {
func (h *Handler) onNodeOnline(nodeID int64) { func (h *Handler) onNodeOnline(nodeID int64) {
h.consumeNodePendingUpgradeRedeploy(nodeID) 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) h.redeployNodeRuntimeAfterUpgrade(nodeID)
} }
@@ -415,7 +405,6 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
return return
} }
// First pass: deploy everything
tunnelFailed := make(map[int64]struct{}) tunnelFailed := make(map[int64]struct{})
for _, tunnelID := range tunnelIDs { for _, tunnelID := range tunnelIDs {
if err := h.redeployTunnelAndForwards(tunnelID); err != nil { if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
@@ -424,9 +413,6 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
} }
} }
// Collect forwards that failed independently (not skipped due to tunnel failure)
var failedForwards []failedForward
for _, forwardID := range forwardIDs { for _, forwardID := range forwardIDs {
forward, getErr := h.getForwardRecord(forwardID) forward, getErr := h.getForwardRecord(forwardID)
if getErr != nil || forward == nil { if getErr != nil || forward == nil {
@@ -436,86 +422,7 @@ func (h *Handler) redeployNodeRuntimeAfterUpgrade(nodeID int64) {
continue continue
} }
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil { 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) 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)
}
} }
+21 -24
View File
@@ -30,22 +30,21 @@ func (User) TableName() string { return "user" }
// Forward maps to the "forward" table. // Forward maps to the "forward" table.
type Forward struct { type Forward struct {
ID int64 `gorm:"primaryKey;autoIncrement"` ID int64 `gorm:"primaryKey;autoIncrement"`
UserID int64 `gorm:"column:user_id;not null"` UserID int64 `gorm:"column:user_id;not null"`
UserName string `gorm:"column:user_name;type:varchar(100);not null"` UserName string `gorm:"column:user_name;type:varchar(100);not null"`
Name string `gorm:"type:varchar(100);not null"` Name string `gorm:"type:varchar(100);not null"`
TunnelID int64 `gorm:"column:tunnel_id;not null"` TunnelID int64 `gorm:"column:tunnel_id;not null"`
RemoteAddr string `gorm:"column:remote_addr;type:text;not null"` RemoteAddr string `gorm:"column:remote_addr;type:text;not null"`
Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"` Strategy string `gorm:"type:varchar(100);not null;default:'fifo'"`
InFlow int64 `gorm:"not null;default:0"` InFlow int64 `gorm:"not null;default:0"`
OutFlow int64 `gorm:"column:out_flow;not null;default:0"` OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
CreatedTime int64 `gorm:"column:created_time;not null"` CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime int64 `gorm:"column:updated_time;not null"` UpdatedTime int64 `gorm:"column:updated_time;not null"`
Status int `gorm:"not null"` Status int `gorm:"not null"`
Inx int `gorm:"not null;default:0"` Inx int `gorm:"not null;default:0"`
SpeedID sql.NullInt64 `gorm:"column:speed_id"` SpeedID sql.NullInt64 `gorm:"column:speed_id"`
MaxConn int `gorm:"column:max_conn;not null;default:0"` 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" } func (Forward) TableName() string { return "forward" }
@@ -440,10 +439,9 @@ type ForwardBackup struct {
CreatedTime int64 `json:"createdTime"` CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"` UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"` Status int `json:"status"`
Inx int `json:"inx"` Inx int `json:"inx"`
SpeedID *int64 `json:"speedId,omitempty"` SpeedID *int64 `json:"speedId,omitempty"`
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"` ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
ProxyProtocol int `json:"proxyProtocol"`
} }
type ForwardPortBackup struct { type ForwardPortBackup struct {
@@ -539,10 +537,9 @@ type ForwardRecord struct {
TunnelID int64 TunnelID int64
RemoteAddr string RemoteAddr string
Strategy string Strategy string
Status int Status int
SpeedID sql.NullInt64 SpeedID sql.NullInt64
MaxConn int MaxConn int
ProxyProtocol int
} }
// TunnelRecord is a minimal tunnel view used by control plane. // TunnelRecord is a minimal tunnel view used by control plane.
+2 -20
View File
@@ -295,17 +295,6 @@ 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 return nil
} }
@@ -725,7 +714,6 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime, "flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
"updatedTime": nullableInt64(u.UpdatedTime), "updatedTime": nullableInt64(u.UpdatedTime),
"inFlow": u.InFlow, "outFlow": u.OutFlow, "inFlow": u.InFlow, "outFlow": u.OutFlow,
"maxConn": u.MaxConn,
} }
if quota := quotaMap[u.ID]; quota != nil { if quota := quotaMap[u.ID]; quota != nil {
item["dailyQuotaGB"] = quota.DailyLimitGB item["dailyQuotaGB"] = quota.DailyLimitGB
@@ -781,13 +769,11 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
Status int Status int
Inx int Inx int
SpeedID sql.NullInt64 SpeedID sql.NullInt64
MaxConn int
ProxyProtocol int
} }
var rows []fwdRow var rows []fwdRow
err := r.db.Model(&model.Forward{}). err := r.db.Model(&model.Forward{}).
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id, forward.max_conn, forward.proxy_protocol"). Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, COALESCE(tunnel.traffic_ratio, 1.0) AS traffic_ratio, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx, forward.speed_id").
Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id"). Joins("LEFT JOIN tunnel ON tunnel.id = forward.tunnel_id").
Order("forward.inx ASC, forward.id ASC"). Order("forward.inx ASC, forward.id ASC").
Find(&rows).Error Find(&rows).Error
@@ -809,8 +795,6 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy, "remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
"inFlow": row.InFlow, "outFlow": row.OutFlow, "inFlow": row.InFlow, "outFlow": row.OutFlow,
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx), "createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
"maxConn": row.MaxConn,
"proxyProtocol": row.ProxyProtocol,
} }
if row.SpeedID.Valid { if row.SpeedID.Valid {
item["speedId"] = row.SpeedID.Int64 item["speedId"] = row.SpeedID.Int64
@@ -2003,7 +1987,6 @@ func (r *Repository) exportForwards() ([]model.ForwardBackup, error) {
TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy, TunnelID: f.TunnelID, RemoteAddr: f.RemoteAddr, Strategy: f.Strategy,
InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime, InFlow: f.InFlow, OutFlow: f.OutFlow, CreatedTime: f.CreatedTime,
UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx, UpdatedTime: f.UpdatedTime, Status: f.Status, Inx: f.Inx,
ProxyProtocol: f.ProxyProtocol,
} }
ports, err := r.exportForwardPorts(f.ID) ports, err := r.exportForwardPorts(f.ID)
if err != nil { if err != nil {
@@ -2401,13 +2384,12 @@ func importForwards(tx *gorm.DB, forwards []model.ForwardBackup, now int64) (int
UpdatedTime: now, UpdatedTime: now,
Status: f.Status, Status: f.Status,
Inx: f.Inx, Inx: f.Inx,
ProxyProtocol: f.ProxyProtocol,
} }
err := tx.Clauses(clause.OnConflict{ err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}}, Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{ DoUpdates: clause.AssignmentColumns([]string{
"user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy", "user_id", "user_name", "name", "tunnel_id", "remote_addr", "strategy",
"in_flow", "out_flow", "updated_time", "status", "inx", "proxy_protocol", "in_flow", "out_flow", "updated_time", "status", "inx",
}), }),
}).Create(&item).Error }).Create(&item).Error
if err != nil { if err != nil {
@@ -124,17 +124,16 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
return nil, err return nil, err
} }
fr := model.ForwardRecord{ fr := model.ForwardRecord{
ID: f.ID, ID: f.ID,
UserID: f.UserID, UserID: f.UserID,
UserName: f.UserName, UserName: f.UserName,
Name: f.Name, Name: f.Name,
TunnelID: f.TunnelID, TunnelID: f.TunnelID,
RemoteAddr: f.RemoteAddr, RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy, Strategy: f.Strategy,
Status: f.Status, Status: f.Status,
SpeedID: f.SpeedID, SpeedID: f.SpeedID,
MaxConn: f.MaxConn, MaxConn: f.MaxConn,
ProxyProtocol: f.ProxyProtocol,
} }
if strings.TrimSpace(fr.Strategy) == "" { if strings.TrimSpace(fr.Strategy) == "" {
fr.Strategy = "fifo" fr.Strategy = "fifo"
@@ -1,59 +0,0 @@
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,21 +695,20 @@ func (r *Repository) GetMinForwardPort(forwardID int64) sql.NullInt64 {
return p return p
} }
func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int, proxyProtocol int) error { func (r *Repository) UpdateForward(id int64, name string, tunnelID int64, remoteAddr, strategy string, now int64, speedID interface{}, maxConn int) error {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return errors.New("repository not initialized") return errors.New("repository not initialized")
} }
return r.db.Model(&model.Forward{}). return r.db.Model(&model.Forward{}).
Where("id = ?", id). Where("id = ?", id).
Updates(map[string]interface{}{ Updates(map[string]interface{}{
"name": name, "name": name,
"tunnel_id": tunnelID, "tunnel_id": tunnelID,
"remote_addr": remoteAddr, "remote_addr": remoteAddr,
"strategy": strategy, "strategy": strategy,
"speed_id": nullInt64FromInterface(speedID), "speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn, "max_conn": maxConn,
"proxy_protocol": proxyProtocol, "updated_time": now,
"updated_time": now,
}).Error }).Error
} }
@@ -783,7 +782,7 @@ func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int,
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error 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, proxyProtocol int, now int64) { func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, maxConn int, now int64) {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return return
} }
@@ -799,7 +798,6 @@ func (r *Repository) RollbackForwardFields(id, userID int64, userName, name stri
"status": status, "status": status,
"speed_id": nullInt64FromInterface(speedID), "speed_id": nullInt64FromInterface(speedID),
"max_conn": maxConn, "max_conn": maxConn,
"proxy_protocol": proxyProtocol,
"updated_time": now, "updated_time": now,
}).Error }).Error
} }
@@ -1260,28 +1258,27 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
return ut.ID, true, nil 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, proxyProtocol int) (int64, error) { func (r *Repository) CreateForwardTx(userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, now int64, inx int, entryNodeIDs []int64, port int, inIp string, speedID interface{}, maxConn int) (int64, error) {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return 0, errors.New("repository not initialized") return 0, errors.New("repository not initialized")
} }
var forwardID int64 var forwardID int64
err := r.db.Transaction(func(tx *gorm.DB) error { err := r.db.Transaction(func(tx *gorm.DB) error {
fwd := model.Forward{ fwd := model.Forward{
UserID: userID, UserID: userID,
UserName: userName, UserName: userName,
Name: name, Name: name,
TunnelID: tunnelID, TunnelID: tunnelID,
RemoteAddr: remoteAddr, RemoteAddr: remoteAddr,
Strategy: strategy, Strategy: strategy,
InFlow: 0, InFlow: 0,
OutFlow: 0, OutFlow: 0,
CreatedTime: now, CreatedTime: now,
UpdatedTime: now, UpdatedTime: now,
Status: 1, Status: 1,
Inx: inx, Inx: inx,
MaxConn: maxConn, MaxConn: maxConn,
SpeedID: nullInt64FromInterface(speedID), SpeedID: nullInt64FromInterface(speedID),
ProxyProtocol: proxyProtocol,
} }
if err := tx.Create(&fwd).Error; err != nil { if err := tx.Create(&fwd).Error; err != nil {
return err return err
@@ -96,7 +96,6 @@ func TestMaxConnLimit(t *testing.T) {
"remoteAddr": "1.1.1.1:443", "remoteAddr": "1.1.1.1:443",
"strategy": "fifo", "strategy": "fifo",
"maxConn": 42, "maxConn": 42,
"proxyProtocol": 2,
} }
body, err := json.Marshal(payload) body, err := json.Marshal(payload)
if err != nil { if err != nil {
@@ -121,47 +120,6 @@ func TestMaxConnLimit(t *testing.T) {
t.Fatalf("get forward ID: %v", err) 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() commandMu.Lock()
defer commandMu.Unlock() defer commandMu.Unlock()
@@ -349,9 +349,9 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
tunnelID := mustLastInsertID(t, r, "backup-forward-tunnel") tunnelID := mustLastInsertID(t, r, "backup-forward-tunnel")
if err := r.DB().Exec(` 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, proxy_protocol) INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88, 2).Error; err != nil { `, 1, "admin_user", "backup-forward", tunnelID, "127.0.0.1:9000", "fifo", 0, 0, now, now, 1, 88).Error; err != nil {
t.Fatalf("seed forward for backup: %v", err) t.Fatalf("seed forward for backup: %v", err)
} }
forwardID := mustLastInsertID(t, r, "backup-forward") forwardID := mustLastInsertID(t, r, "backup-forward")
@@ -412,9 +412,6 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
if !ok { if !ok {
t.Fatalf("expected forwardPorts for forward %d in payload", forwardID) 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 { for _, p := range portsRaw {
portMap, ok := p.(map[string]interface{}) portMap, ok := p.(map[string]interface{})
if !ok { if !ok {
@@ -478,14 +475,6 @@ func TestBackupExportImportRestoreContracts(t *testing.T) {
t.Fatalf("expected forward_port node=%d port=%d after import, got %v", nodeID, port, after) 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) { t.Run("backup export tolerates nullable legacy tunnel chain fields", func(t *testing.T) {
@@ -1,59 +0,0 @@
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)
}
}
-4
View File
@@ -116,10 +116,6 @@ func main() {
fmt.Printf("✅ 配置加载成功 - addr: %s\n", config.Addr) fmt.Printf("✅ 配置加载成功 - addr: %s\n", config.Addr)
// 设置运行时配置持久化路径
socket.SetConfigPersistPath("gost.json")
// 启用持久化将在 program.Start() 后开启,避免启动加载阶段触发冗余写入
log := xlogger.NewLogger() log := xlogger.NewLogger()
logger.SetDefault(log) logger.SetDefault(log)
-5
View File
@@ -16,7 +16,6 @@ import (
metrics "github.com/go-gost/x/metrics/service" metrics "github.com/go-gost/x/metrics/service"
"github.com/go-gost/x/registry" "github.com/go-gost/x/registry"
xservice "github.com/go-gost/x/service" xservice "github.com/go-gost/x/service"
"github.com/go-gost/x/socket"
"github.com/judwhite/go-svc" "github.com/judwhite/go-svc"
"net/http" "net/http"
"os" "os"
@@ -67,10 +66,6 @@ func (p *program) Start() error {
return err 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 { if err := p.run(cfg); err != nil {
return err return err
} }
+2 -9
View File
@@ -44,14 +44,9 @@ func Set(c *Config) {
func OnUpdate(f func(c *Config) error) error { func OnUpdate(f func(c *Config) error) error {
globalMux.Lock() globalMux.Lock()
err := f(global) defer globalMux.Unlock()
globalMux.Unlock()
if err == nil { return f(global)
persist()
}
return err
} }
type LogConfig struct { type LogConfig struct {
@@ -578,7 +573,6 @@ func (c *Config) Load() error {
if err := v.ReadInConfig(); err != nil { if err := v.ReadInConfig(); err != nil {
return err return err
} }
SetPersistPath(v.ConfigFileUsed())
return v.Unmarshal(c) return v.Unmarshal(c)
} }
@@ -596,7 +590,6 @@ func (c *Config) ReadFile(file string) error {
if err := v.ReadInConfig(); err != nil { if err := v.ReadInConfig(); err != nil {
return err return err
} }
SetPersistPath(v.ConfigFileUsed())
return v.Unmarshal(c) return v.Unmarshal(c)
} }
-94
View File
@@ -1,94 +0,0 @@
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)
}
-32
View File
@@ -1,32 +0,0 @@
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)
}
}
-20
View File
@@ -1690,26 +1690,6 @@ func StartWebSocketReporterWithConfig(addr string, secret string, http int, tls
return reporter 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诊断命令 // handleTcpPing 处理TCP ping诊断命令
func (w *WebSocketReporter) handleTcpPing(data interface{}) (TcpPingResponse, error) { func (w *WebSocketReporter) handleTcpPing(data interface{}) (TcpPingResponse, error) {
jsonData, err := json.Marshal(data) jsonData, err := json.Marshal(data)
+27 -18
View File
@@ -5,7 +5,7 @@ import {
useNavigate, useNavigate,
Navigate, Navigate,
} from "react-router-dom"; } from "react-router-dom";
import { useEffect } from "react"; import { useEffect, useState } from "react";
import { AnimatePresence } from "framer-motion"; import { AnimatePresence } from "framer-motion";
import IndexPage from "@/pages/index"; import IndexPage from "@/pages/index";
@@ -39,7 +39,32 @@ const ProtectedRoute = ({
skipLayout?: boolean; skipLayout?: boolean;
}) => { }) => {
const isH5 = useH5Mode(); const isH5 = useH5Mode();
const authenticated = isLoggedIn(); 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]);
if (!authenticated) { if (!authenticated) {
return <Navigate replace to="/" />; return <Navigate replace to="/" />;
@@ -78,22 +103,6 @@ const LoginRoute = () => {
function App() { function App() {
const location = useLocation(); 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(() => { useEffect(() => {
+1 -5
View File
@@ -71,8 +71,6 @@ export interface ForwardApiItem {
userId?: number; userId?: number;
tunnelId?: number; tunnelId?: number;
speedId?: number | null; speedId?: number | null;
maxConn?: number;
proxyProtocol?: number;
inx?: number; inx?: number;
[key: string]: unknown; [key: string]: unknown;
} }
@@ -203,7 +201,7 @@ export interface UserPackageInfoApiData {
num: number; num: number;
expTime?: string; expTime?: string;
flowResetTime?: number; flowResetTime?: number;
maxConn?: number; maxConn?: number;
[key: string]: unknown; [key: string]: unknown;
}; };
tunnelPermissions: UserTunnelPermissionApiItem[]; tunnelPermissions: UserTunnelPermissionApiItem[];
@@ -372,8 +370,6 @@ export interface ForwardMutationPayload {
remoteAddr?: string; remoteAddr?: string;
strategy?: string; strategy?: string;
speedId?: number | null; speedId?: number | null;
maxConn?: number;
proxyProtocol?: number;
} }
export interface SpeedLimitMutationPayload { export interface SpeedLimitMutationPayload {
+1
View File
@@ -241,6 +241,7 @@ export default function AdminLayout({
// 退出登录 // 退出登录
const handleLogout = () => { const handleLogout = () => {
safeLogout(); safeLogout();
navigate("/");
}; };
// 切换移动端菜单 // 切换移动端菜单
@@ -1,4 +1,5 @@
import { useState } from "react"; import { useState } from "react";
import { useNavigate } from "react-router-dom";
import toast from "react-hot-toast"; import toast from "react-hot-toast";
import { Button } from "@/shadcn-bridge/heroui/button"; import { Button } from "@/shadcn-bridge/heroui/button";
@@ -25,6 +26,7 @@ export default function ChangePasswordPage() {
}); });
const [loading, setLoading] = useState(false); const [loading, setLoading] = useState(false);
const [errors, setErrors] = useState<Partial<PasswordForm>>({}); const [errors, setErrors] = useState<Partial<PasswordForm>>({});
const navigate = useNavigate();
const validateForm = (): boolean => { const validateForm = (): boolean => {
const newErrors: Partial<PasswordForm> = {}; const newErrors: Partial<PasswordForm> = {};
@@ -96,6 +98,7 @@ export default function ChangePasswordPage() {
const logout = () => { const logout = () => {
safeLogout(); safeLogout();
navigate("/");
}; };
const handleKeyPress = (e: React.KeyboardEvent) => { const handleKeyPress = (e: React.KeyboardEvent) => {
+1
View File
@@ -50,6 +50,7 @@ export default function DashboardPage() {
const handleLogout = () => { const handleLogout = () => {
safeLogout(); safeLogout();
toast.success("已退出登录"); toast.success("已退出登录");
navigate("/");
}; };
const { const {
+13 -53
View File
@@ -125,7 +125,6 @@ interface Forward {
userId?: number; userId?: number;
inx?: number; inx?: number;
speedId?: number | null; speedId?: number | null;
proxyProtocol?: number;
} }
interface Tunnel { interface Tunnel {
@@ -161,7 +160,6 @@ interface ForwardForm {
strategy: string; strategy: string;
speedId: number | null; speedId: number | null;
maxConn?: number; maxConn?: number;
proxyProtocol?: number;
} }
interface ForwardUserGroup { interface ForwardUserGroup {
@@ -578,11 +576,6 @@ const mapForwardApiItems = (items: ForwardApiItem[]): Forward[] => {
typeof forward.speedId === "number" || forward.speedId === null typeof forward.speedId === "number" || forward.speedId === null
? forward.speedId ? forward.speedId
: undefined, : undefined,
maxConn: typeof forward.maxConn === "number" ? forward.maxConn : undefined,
proxyProtocol:
typeof forward.proxyProtocol === "number"
? forward.proxyProtocol
: undefined,
serviceRunning: forward.status === 1, serviceRunning: forward.status === 1,
})); }));
}; };
@@ -1317,7 +1310,6 @@ export default function ForwardPage() {
strategy: "fifo", strategy: "fifo",
speedId: null, speedId: null,
maxConn: 0, maxConn: 0,
proxyProtocol: 0,
}); });
const [inIpTouched, setInIpTouched] = useState(false); const [inIpTouched, setInIpTouched] = useState(false);
@@ -2105,7 +2097,6 @@ export default function ForwardPage() {
interfaceName: "", interfaceName: "",
strategy: "fifo", strategy: "fifo",
speedId: null, speedId: null,
proxyProtocol: 0,
}); });
setErrors({}); setErrors({});
setModalOpen(true); setModalOpen(true);
@@ -2126,8 +2117,7 @@ export default function ForwardPage() {
interfaceName: forward.interfaceName || "", interfaceName: forward.interfaceName || "",
strategy: forward.strategy || "fifo", strategy: forward.strategy || "fifo",
speedId: normalizeSpeedId(forward.speedId), speedId: normalizeSpeedId(forward.speedId),
maxConn: forward.maxConn ?? 0, maxConn: forward.maxConn || 0,
proxyProtocol: forward.proxyProtocol ?? 0,
}); });
setErrors({}); setErrors({});
setModalOpen(true); setModalOpen(true);
@@ -2257,7 +2247,6 @@ export default function ForwardPage() {
strategy: addressCount > 1 ? form.strategy : "fifo", strategy: addressCount > 1 ? form.strategy : "fifo",
speedId: normalizedSpeedId, speedId: normalizedSpeedId,
maxConn: form.maxConn, maxConn: form.maxConn,
proxyProtocol: form.proxyProtocol,
}; };
res = await updateForward(updateData); res = await updateForward(updateData);
@@ -2271,11 +2260,11 @@ export default function ForwardPage() {
strategy: addressCount > 1 ? form.strategy : "fifo", strategy: addressCount > 1 ? form.strategy : "fifo",
speedId: normalizedSpeedId, speedId: normalizedSpeedId,
maxConn: form.maxConn, maxConn: form.maxConn,
proxyProtocol: form.proxyProtocol,
}; };
res = await createForward(createData); res = await createForward(createData);
} }
if (res.code === 0) { if (res.code === 0) {
const warningItems = Array.isArray((res as any).data?.warnings) const warningItems = Array.isArray((res as any).data?.warnings)
? (res as any).data.warnings ? (res as any).data.warnings
@@ -4753,6 +4742,8 @@ export default function ForwardPage() {
} }
/> />
<Select <Select
description={ description={
isEdit isEdit
@@ -4880,55 +4871,26 @@ export default function ForwardPage() {
<SelectItem key="hash">哈希模式 - IP哈希</SelectItem> <SelectItem key="hash">哈希模式 - IP哈希</SelectItem>
</Select> </Select>
)} )}
<Accordion className="px-0" variant="light"> <Accordion className="px-0" variant="light">
<AccordionItem <AccordionItem
key="advanced" key="advanced"
aria-label="高级设置" aria-label="高级设置"
title={ title={<span className="text-small text-default-500 font-medium">高级设置</span>}
<span className="text-small text-default-500 font-medium">
高级设置
</span>
}
> >
<div className="space-y-4 pb-2"> <div className="space-y-4 pb-2">
<Input <Input
description="此设置优先于用户的全局连接数限制。0 表示不限制。"
label="最大连接数" label="最大连接数"
min="0"
placeholder="0 或空表示不限制" placeholder="0 或空表示不限制"
type="number" type="number"
value={ min="0"
form.maxConn === 0 ? "" : String(form.maxConn || "") value={form.maxConn === 0 ? "" : String(form.maxConn || "")}
}
variant="bordered"
onChange={(e) => { onChange={(e) => {
const value = Math.max( const value = Math.max(Number(e.target.value) || 0, 0);
Number(e.target.value) || 0,
0,
);
setForm((prev) => ({ ...prev, maxConn: value })); setForm((prev) => ({ ...prev, maxConn: value }));
}} }}
/> description="此设置优先于用户的全局连接数限制。0 表示不限制。"
<Select
description="启用 PROXY protocol,用于透传客户端真实 IP"
label="Proxy Protocol"
placeholder="禁用"
selectedKeys={[String(form.proxyProtocol || 0)]}
variant="bordered" 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 && ( {isAdmin && (
<Select <Select
label="规则限速" label="规则限速"
@@ -4946,9 +4908,7 @@ export default function ForwardPage() {
setForm((prev) => ({ setForm((prev) => ({
...prev, ...prev,
speedId: selectedKey speedId: selectedKey ? Number(selectedKey) : null,
? Number(selectedKey)
: null,
})); }));
}} }}
> >
@@ -4965,7 +4925,7 @@ export default function ForwardPage() {
</div> </div>
</AccordionItem> </AccordionItem>
</Accordion> </Accordion>
</div> </div>
</ModalBody> </ModalBody>
<ModalFooter> <ModalFooter>
<Button variant="light" onPress={onClose}> <Button variant="light" onPress={onClose}>
+1
View File
@@ -127,6 +127,7 @@ export default function ProfilePage() {
// 退出登录 // 退出登录
const handleLogout = () => { const handleLogout = () => {
safeLogout(); safeLogout();
navigate("/", { replace: true });
}; };
// 密码表单验证 // 密码表单验证
+3 -8
View File
@@ -160,7 +160,6 @@ const normalizeUserItem = (item: Partial<User>): User => {
monthlyUsedBytes: Number(item.monthlyUsedBytes ?? 0), monthlyUsedBytes: Number(item.monthlyUsedBytes ?? 0),
disabledByQuota: Number(item.disabledByQuota ?? 0), disabledByQuota: Number(item.disabledByQuota ?? 0),
quotaDisabledAt: Number(item.quotaDisabledAt ?? 0), quotaDisabledAt: Number(item.quotaDisabledAt ?? 0),
maxConn: item.maxConn != null ? Number(item.maxConn) : undefined,
}; };
}; };
@@ -516,7 +515,7 @@ export default function UserPage() {
num: 10, num: 10,
expTime: null, expTime: null,
flowResetTime: 0, flowResetTime: 0,
maxConn: 0, maxConn: 0,
groupIds: [], groupIds: [],
}); });
onUserModalOpen(); onUserModalOpen();
@@ -546,7 +545,6 @@ export default function UserPage() {
num: user.num, num: user.num,
expTime: user.expTime ? new Date(user.expTime) : null, expTime: user.expTime ? new Date(user.expTime) : null,
flowResetTime: user.flowResetTime ?? 0, flowResetTime: user.flowResetTime ?? 0,
maxConn: user.maxConn ?? 0,
groupIds: currentGroupIds, groupIds: currentGroupIds,
}); });
onUserModalOpen(); onUserModalOpen();
@@ -1522,15 +1520,12 @@ export default function UserPage() {
/> />
<Input <Input
label="最大连接数" label="最大连接数"
min="0"
placeholder="0 或空表示不限制" placeholder="0 或空表示不限制"
type="number" type="number"
value={ min="0"
userForm.maxConn === 0 ? "" : String(userForm.maxConn || "") value={userForm.maxConn === 0 ? "" : String(userForm.maxConn || "")}
}
onChange={(e) => { onChange={(e) => {
const value = Math.max(Number(e.target.value) || 0, 0); const value = Math.max(Number(e.target.value) || 0, 0);
setUserForm((prev) => ({ ...prev, maxConn: value })); setUserForm((prev) => ({ ...prev, maxConn: value }));
}} }}
/> />
-1
View File
@@ -24,7 +24,6 @@ export interface User {
monthlyUsedBytes?: number; monthlyUsedBytes?: number;
disabledByQuota?: number; disabledByQuota?: number;
quotaDisabledAt?: number; quotaDisabledAt?: number;
maxConn?: number;
} }
export interface UserGroup { export interface UserGroup {
+2 -1
View File
@@ -2,8 +2,9 @@ import { clearSession } from "@/utils/session";
/** /**
* 安全退出登录函数 * 安全退出登录函数
* 清除登录相关数据 * 清除登录相关数据,并强制刷新跳转到首页
*/ */
export const safeLogout = () => { export const safeLogout = () => {
clearSession(); clearSession();
window.location.href = "/";
}; };