feat: decouple speed limits from tunnels and add forward-level rate limiting

- Make SpeedLimit.TunnelID and TunnelName nullable (optional binding)
- Add SpeedID field to Forward model for forward-level rate limiting
- Update ForwardRecord to include SpeedID for control plane
- Update repository methods to handle optional tunnel binding
- Update handlers to accept optional tunnelId in create/update
- Modify control plane to prioritize Forward.SpeedID over UserTunnel speed limit
- Update frontend limit.tsx to support creating speed limits without tunnel binding
- Update TypeScript types for optional tunnelId and new speedId fields

This allows speed limits to be created as reusable rules that can be applied
to either tunnels (via UserTunnel.SpeedID) or individual forwards (via Forward.SpeedID).
This commit is contained in:
sagitchu
2026-02-26 13:08:35 +08:00
parent 4bdfa50b0c
commit c8eb780c67
11 changed files with 318 additions and 89 deletions
@@ -152,11 +152,32 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
return errors.New("转发入口端口不存在")
}
userTunnelID, limiterID, speed, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil {
return err
// Determine limiter from forward's SpeedID first, fallback to UserTunnel's limiter
var limiterID *int64
var speed *int
if forward.SpeedID.Valid && forward.SpeedID.Int64 > 0 {
// Forward has its own speed limit
speedVal, err := h.repo.GetSpeedLimitSpeed(forward.SpeedID.Int64)
if err == nil && speedVal > 0 {
limiterID = &forward.SpeedID.Int64
speed = &speedVal
}
}
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, userTunnelID)
if limiterID == nil {
// Fall back to UserTunnel speed limit
var utLimiterID *int64
var utSpeed *int
_, utLimiterID, utSpeed, err = h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
if err != nil {
return err
}
limiterID = utLimiterID
speed = utSpeed
}
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, 0)
tunnelTLSProtocol, err := h.isTunnelSelectedTLSProtocol(forward.TunnelID)
if err != nil {
return err
+51 -22
View File
@@ -1586,29 +1586,37 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
tunnelID := asInt64(req["tunnelId"], 0)
if tunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
name := asString(req["name"])
if name == "" {
response.WriteJSON(w, response.ErrDefault("名称不能为空"))
return
}
tunnelName := h.repo.GetTunnelNameByID(tunnelID)
if tunnelName == "" {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
now := time.Now().UnixMilli()
speed := asInt(req["speed"], 100)
var tunnelID *int64
var tunnelName string
if tid := asInt64(req["tunnelId"], 0); tid > 0 {
tunnelID = &tid
tunnelName = h.repo.GetTunnelNameByID(tid)
if tunnelName == "" {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
}
now := time.Now().UnixMilli()
id, err := h.repo.CreateSpeedLimit(name, speed, tunnelID, tunnelName, now, asInt(req["status"], 1))
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.sendLimiterConfig(id, speed, tunnelID)
if tunnelID != nil && *tunnelID > 0 {
_ = h.sendLimiterConfig(id, speed, *tunnelID)
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -1618,23 +1626,41 @@ func (h *Handler) speedLimitUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
id := asInt64(req["id"], 0)
tunnelID := asInt64(req["tunnelId"], 0)
if id <= 0 || tunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
if id <= 0 {
response.WriteJSON(w, response.ErrDefault("限速规则ID不能为空"))
return
}
tunnelName := h.repo.GetTunnelNameByID(tunnelID)
if tunnelName == "" {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
name := asString(req["name"])
if name == "" {
response.WriteJSON(w, response.ErrDefault("名称不能为空"))
return
}
speed := asInt(req["speed"], 100)
if err := h.repo.UpdateSpeedLimit(id, asString(req["name"]), speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil {
var tunnelID *int64
var tunnelName string
if tid := asInt64(req["tunnelId"], 0); tid > 0 {
tunnelID = &tid
tunnelName = h.repo.GetTunnelNameByID(tid)
if tunnelName == "" {
response.WriteJSON(w, response.ErrDefault("隧道不存在"))
return
}
}
if err := h.repo.UpdateSpeedLimit(id, name, speed, tunnelID, tunnelName, asInt(req["status"], 1), time.Now().UnixMilli()); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
_ = h.sendLimiterConfig(id, speed, tunnelID)
if tunnelID != nil && *tunnelID > 0 {
_ = h.sendLimiterConfig(id, speed, *tunnelID)
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -1643,15 +1669,18 @@ func (h *Handler) speedLimitDelete(w http.ResponseWriter, r *http.Request) {
if id <= 0 {
return
}
tunnelID := h.repo.GetSpeedLimitTunnelID(id)
if err := h.repo.DeleteSpeedLimit(id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if tunnelID > 0 {
_ = h.sendDeleteLimiterConfig(id, tunnelID)
if tunnelID.Valid && tunnelID.Int64 > 0 {
_ = h.sendDeleteLimiterConfig(id, tunnelID.Int64)
}
response.WriteJSON(w, response.OKEmpty())
}