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())
}
+26 -23
View File
@@ -29,19 +29,20 @@ 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:"column:in_flow;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"`
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"`
}
func (Forward) TableName() string { return "forward" }
@@ -83,14 +84,14 @@ type Node struct {
func (Node) TableName() string { return "node" }
type SpeedLimit struct {
ID int64 `gorm:"primaryKey;autoIncrement"`
Name string `gorm:"type:varchar(100);not null"`
Speed int `gorm:"not null"`
TunnelID int64 `gorm:"column:tunnel_id;not null"`
TunnelName string `gorm:"column:tunnel_name;type:varchar(100);not null"`
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
Status int `gorm:"not null"`
ID int64 `gorm:"primaryKey;autoIncrement"`
Name string `gorm:"type:varchar(100);not null"`
Speed int `gorm:"not null"`
TunnelID sql.NullInt64 `gorm:"column:tunnel_id"`
TunnelName sql.NullString `gorm:"column:tunnel_name;type:varchar(100)"`
CreatedTime int64 `gorm:"column:created_time;not null"`
UpdatedTime sql.NullInt64 `gorm:"column:updated_time"`
Status int `gorm:"not null"`
}
func (SpeedLimit) TableName() string { return "speed_limit" }
@@ -395,6 +396,7 @@ type ForwardBackup struct {
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
Inx int `json:"inx"`
SpeedID *int64 `json:"speedId,omitempty"`
ForwardPorts *[]ForwardPortBackup `json:"forwardPorts,omitempty"`
}
@@ -421,8 +423,8 @@ type SpeedLimitBackup struct {
ID int64 `json:"id"`
Name string `json:"name"`
Speed int64 `json:"speed"`
TunnelID int64 `json:"tunnelId"`
TunnelName string `json:"tunnelName"`
TunnelID *int64 `json:"tunnelId,omitempty"`
TunnelName string `json:"tunnelName,omitempty"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime,omitempty"`
Status int `json:"status"`
@@ -492,6 +494,7 @@ type ForwardRecord struct {
RemoteAddr string
Strategy string
Status int
SpeedID sql.NullInt64
}
// TunnelRecord is a minimal tunnel view used by control plane.
+32 -9
View File
@@ -683,12 +683,18 @@ func (r *Repository) ListSpeedLimits() ([]map[string]interface{}, error) {
}
items := make([]map[string]interface{}, 0, len(limits))
for _, sl := range limits {
items = append(items, map[string]interface{}{
item := map[string]interface{}{
"id": sl.ID, "name": sl.Name, "speed": sl.Speed,
"tunnelId": sl.TunnelID, "tunnelName": sl.TunnelName,
"status": sl.Status, "createdTime": sl.CreatedTime,
"updatedTime": nullableInt64(sl.UpdatedTime),
})
}
if sl.TunnelID.Valid {
item["tunnelId"] = sl.TunnelID.Int64
}
if sl.TunnelName.Valid {
item["tunnelName"] = sl.TunnelName.String
}
items = append(items, item)
}
return items, nil
}
@@ -712,11 +718,12 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
CreatedTime int64
Status int
Inx int
SpeedID sql.NullInt64
}
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, forward.remote_addr, COALESCE(forward.strategy, 'fifo') AS strategy, forward.in_flow, forward.out_flow, forward.created_time, forward.status, forward.inx").
Select("forward.id, forward.user_id, forward.user_name, forward.name, forward.tunnel_id, COALESCE(tunnel.name, '') AS tunnel_name, 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").
Order("forward.inx ASC, forward.id ASC").
Find(&rows).Error
@@ -730,14 +737,18 @@ func (r *Repository) ListForwards() ([]map[string]interface{}, error) {
if err != nil {
return nil, err
}
items = append(items, map[string]interface{}{
item := map[string]interface{}{
"id": row.ID, "userId": row.UserID, "userName": row.UserName,
"name": row.Name, "tunnelId": row.TunnelID, "tunnelName": row.TunnelName,
"inIp": nullableForwardIngress(inIP), "inPort": nullableInt64(inPort),
"remoteAddr": row.RemoteAddr, "strategy": row.Strategy,
"inFlow": row.InFlow, "outFlow": row.OutFlow,
"createdTime": row.CreatedTime, "status": row.Status, "inx": int64(row.Inx),
})
}
if row.SpeedID.Valid {
item["speedId"] = row.SpeedID.Int64
}
items = append(items, item)
}
return items, nil
}
@@ -1813,9 +1824,15 @@ func (r *Repository) exportSpeedLimits() ([]model.SpeedLimitBackup, error) {
for _, sl := range sls {
b := model.SpeedLimitBackup{
ID: sl.ID, Name: sl.Name, Speed: int64(sl.Speed),
TunnelID: sl.TunnelID, TunnelName: sl.TunnelName,
CreatedTime: sl.CreatedTime, Status: sl.Status,
}
if sl.TunnelID.Valid {
tid := sl.TunnelID.Int64
b.TunnelID = &tid
}
if sl.TunnelName.Valid {
b.TunnelName = sl.TunnelName.String
}
if sl.UpdatedTime.Valid {
b.UpdatedTime = sl.UpdatedTime.Int64
}
@@ -2186,12 +2203,18 @@ func importSpeedLimits(tx *gorm.DB, speedLimits []model.SpeedLimitBackup, now in
ID: sl.ID,
Name: sl.Name,
Speed: int(sl.Speed),
TunnelID: sl.TunnelID,
TunnelName: sl.TunnelName,
TunnelID: sql.NullInt64{Int64: 0, Valid: false},
TunnelName: sql.NullString{String: "", Valid: false},
CreatedTime: sl.CreatedTime,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: sl.Status,
}
if sl.TunnelID != nil {
item.TunnelID = sql.NullInt64{Int64: *sl.TunnelID, Valid: true}
}
if sl.TunnelName != "" {
item.TunnelName = sql.NullString{String: sl.TunnelName, Valid: true}
}
err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "id"}},
DoUpdates: clause.AssignmentColumns([]string{
@@ -46,6 +46,7 @@ func (r *Repository) ListForwardsByTunnel(tunnelID int64) ([]model.ForwardRecord
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
})
}
for i := range rows {
@@ -38,6 +38,7 @@ func (r *Repository) ListActiveForwardsByUser(userID int64) ([]model.ForwardReco
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
})
}
for i := range rows {
@@ -68,6 +69,7 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
})
}
for i := range rows {
@@ -99,6 +101,7 @@ func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, er
RemoteAddr: f.RemoteAddr,
Strategy: f.Strategy,
Status: f.Status,
SpeedID: f.SpeedID,
}
if strings.TrimSpace(fr.Strategy) == "" {
fr.Strategy = "fifo"
@@ -169,3 +172,15 @@ func (r *Repository) SpeedLimitExists(id int64) (bool, error) {
}
return count > 0, nil
}
func (r *Repository) GetSpeedLimitSpeed(id int64) (int, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
var sl model.SpeedLimit
err := r.db.Select("speed").Where("id = ?", id).First(&sl).Error
if err != nil {
return 0, err
}
return sl.Speed, nil
}
@@ -764,51 +764,66 @@ func (r *Repository) GetUsedPortsOnNodeAsMap(nodeID int64) (map[int]bool, error)
return used, nil
}
func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID int64, tunnelName string, now int64, status int) (int64, error) {
func (r *Repository) CreateSpeedLimit(name string, speed int, tunnelID *int64, tunnelName string, now int64, status int) (int64, error) {
if r == nil || r.db == nil {
return 0, errors.New("repository not initialized")
}
sl := model.SpeedLimit{
Name: name,
Speed: speed,
TunnelID: tunnelID,
TunnelName: tunnelName,
TunnelID: sql.NullInt64{Int64: 0, Valid: false},
TunnelName: sql.NullString{String: "", Valid: false},
CreatedTime: now,
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: status,
}
if tunnelID != nil {
sl.TunnelID = sql.NullInt64{Int64: *tunnelID, Valid: true}
}
if tunnelName != "" {
sl.TunnelName = sql.NullString{String: tunnelName, Valid: true}
}
if err := r.db.Create(&sl).Error; err != nil {
return 0, err
}
return sl.ID, nil
}
func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID int64, tunnelName string, status int, now int64) error {
func (r *Repository) UpdateSpeedLimit(id int64, name string, speed int, tunnelID *int64, tunnelName string, status int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
updates := map[string]interface{}{
"name": name,
"speed": speed,
"status": status,
"updated_time": sql.NullInt64{
Int64: now,
Valid: true,
},
}
if tunnelID != nil {
updates["tunnel_id"] = sql.NullInt64{Int64: *tunnelID, Valid: true}
} else {
updates["tunnel_id"] = sql.NullInt64{Int64: 0, Valid: false}
}
if tunnelName != "" {
updates["tunnel_name"] = sql.NullString{String: tunnelName, Valid: true}
} else {
updates["tunnel_name"] = sql.NullString{String: "", Valid: false}
}
return r.db.Model(&model.SpeedLimit{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"name": name,
"speed": speed,
"tunnel_id": tunnelID,
"tunnel_name": tunnelName,
"status": status,
"updated_time": sql.NullInt64{
Int64: now,
Valid: true,
},
}).Error
Updates(updates).Error
}
func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) int64 {
func (r *Repository) GetSpeedLimitTunnelID(speedLimitID int64) sql.NullInt64 {
if r == nil || r.db == nil {
return 0
return sql.NullInt64{Valid: false}
}
var sl model.SpeedLimit
if err := r.db.Select("tunnel_id").Where("id = ?", speedLimitID).First(&sl).Error; err != nil {
return 0
return sql.NullInt64{Valid: false}
}
return sl.TunnelID
}