fix(backend): block forwards when flow exceeded

This commit is contained in:
sagitchu
2026-03-11 09:34:36 +08:00
parent 32ee511eac
commit 2e8c0530a9
4 changed files with 294 additions and 2 deletions
@@ -312,6 +312,14 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
return warnings, fmt.Errorf("节点 %s 下发失败: %w", node.Name, err)
}
}
// Keep paused forwards paused after UpdateService/AddService, since agent-side UpdateService
// always restarts services.
if forward.Status != 1 {
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
return warnings, err
}
}
return warnings, nil
}
@@ -1620,6 +1628,13 @@ func processServerAddress(serverAddr string) string {
if strings.HasPrefix(serverAddr, "[") {
return serverAddr
}
// If the input is a bare IPv6 host (no port), bracket it.
// IPv6-with-port must be provided in bracket form: [::1]:443.
if looksLikeIPv6(serverAddr) {
if ip := net.ParseIP(serverAddr); ip != nil && ip.To4() == nil {
return "[" + serverAddr + "]"
}
}
idx := strings.LastIndex(serverAddr, ":")
if idx < 0 {
@@ -2,6 +2,7 @@ package handler
import (
"encoding/json"
"errors"
"log"
"strconv"
"strings"
@@ -327,6 +328,67 @@ func (h *Handler) enforceFlowPolicies(userID int64, userTunnelID int64) {
}
}
func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, now int64) error {
if h == nil || h.repo == nil {
return errors.New("invalid flow policy context")
}
if userID <= 0 || tunnelID <= 0 {
return nil
}
user, err := h.repo.GetUserByID(userID)
if err != nil {
return err
}
if user == nil {
return errors.New("用户不存在")
}
if user.Status != 1 {
return errors.New("账号已禁用")
}
if user.ExpTime > 0 && user.ExpTime <= now {
return errors.New("账号已过期")
}
flowLimit := user.Flow * bytesPerGB
current := user.InFlow + user.OutFlow
if flowLimit < current {
return errors.New("流量已超额,禁止开启转发")
}
userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(userID, tunnelID)
if err != nil {
return err
}
if userTunnelID <= 0 {
return nil
}
policy, err := h.getUserTunnelPolicy(userTunnelID)
if err != nil {
return err
}
if policy == nil {
return nil
}
if policy.Status != 1 {
return errors.New("该隧道已禁用")
}
if policy.ExpTime > 0 && policy.ExpTime <= now {
return errors.New("该隧道已过期")
}
utFlowLimit := policy.Flow * bytesPerGB
utCurrent := policy.InFlow + policy.OutFlow
if utCurrent >= utFlowLimit {
return errors.New("该隧道流量已超额,禁止开启转发")
}
return nil
}
func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
user, err := h.repo.GetUserByID(userID)
if err != nil || user == nil {
+16 -2
View File
@@ -1155,6 +1155,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法创建转发"))
return
}
if err := h.ensureUserTunnelForwardAllowed(userID, tunnelID, time.Now().UnixMilli()); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
name := asString(req["name"])
remoteAddr := asString(req["remoteAddr"])
if name == "" || remoteAddr == "" {
@@ -1444,11 +1448,16 @@ func (h *Handler) forwardResume(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
now := time.Now().UnixMilli()
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
_ = h.repo.UpdateForwardStatus(id, 1, time.Now().UnixMilli())
_ = h.repo.UpdateForwardStatus(id, 1, now)
response.WriteJSON(w, response.OKEmpty())
}
@@ -1571,17 +1580,22 @@ func (h *Handler) forwardBatchResume(w http.ResponseWriter, r *http.Request) {
}
s := 0
f := 0
now := time.Now().UnixMilli()
for _, id := range ids {
forward, accessErr := h.ensureForwardAccessByActor(actorUserID, actorRole, id)
if accessErr != nil {
f++
continue
}
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
f++
continue
}
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
f++
continue
}
if err := h.repo.UpdateForwardStatus(id, 1, time.Now().UnixMilli()); err != nil {
if err := h.repo.UpdateForwardStatus(id, 1, now); err != nil {
f++
} else {
s++