feat: add speed limit contract tests and refine limit/user UI

This commit is contained in:
sagitchu
2026-02-26 20:18:29 +08:00
parent c8eb780c67
commit 61c5b5e759
12 changed files with 1304 additions and 91 deletions
@@ -185,7 +185,9 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
for _, fp := range ports {
if limiterID != nil && speed != nil {
h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed)
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
return err
}
}
node, err := h.getNodeRecord(fp.NodeID)
@@ -1051,12 +1053,16 @@ func (h *Handler) sendDeleteLimiterConfig(limiterID int64, tunnelID int64) error
return nil
}
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) {
func (h *Handler) ensureLimiterOnNode(nodeID int64, limiterID int64, speed int) error {
rate := float64(speed) / 8.0
limitStr := fmt.Sprintf("$ %.1fMB %.1fMB", rate, rate)
payload := map[string]interface{}{
"name": strconv.FormatInt(limiterID, 10),
"limits": []string{limitStr},
}
_, _ = h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false)
if _, err := h.sendNodeCommand(nodeID, "AddLimiters", payload, false, false); err != nil {
return fmt.Errorf("限速规则下发失败: %w", err)
}
return nil
}
+137 -15
View File
@@ -1048,21 +1048,55 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("权限ID不能为空"))
return
}
speedID := asAnyToInt64Ptr(req["speedId"])
if err := h.validateSpeedLimitReference(speedID); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
if utErr != nil {
response.WriteJSON(w, response.Err(-2, utErr.Error()))
return
}
_, oldFlow, oldNum, oldExpTime, oldFlowReset, oldSpeedID, oldStatus, oldErr :=
h.repo.GetExistingUserTunnel(userID, tunnelID)
if oldErr != nil {
response.WriteJSON(w, response.Err(-2, oldErr.Error()))
return
}
if err := h.repo.UpdateUserTunnel(id,
asInt64(req["flow"], 0),
asInt(req["num"], 0),
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
asInt64(req["flowResetTime"], 1),
nullableInt(asAnyToInt64Ptr(req["speedId"])),
nullableInt(speedID),
asInt(req["status"], 1),
); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
userID, tunnelID, utErr := h.repo.GetUserTunnelUserAndTunnel(id)
if utErr == nil {
h.syncUserTunnelForwards(userID, tunnelID)
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
rollbackErr := h.repo.UpdateUserTunnel(
id,
oldFlow,
int(oldNum),
oldExpTime,
oldFlowReset,
oldSpeedID,
oldStatus,
)
if rollbackErr != nil {
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr)))
return
}
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败,已回滚: %v", syncErr)))
return
}
response.WriteJSON(w, response.OKEmpty())
@@ -1103,6 +1137,18 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.ErrDefault("转发名称和目标地址不能为空"))
return
}
speedID := asAnyToInt64Ptr(req["speedId"])
if speedID != nil {
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
if speedErr != nil {
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
return
}
if !exists {
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
return
}
}
port := asInt(req["inPort"], 0)
if port <= 0 {
port = h.pickTunnelPort(tunnelID)
@@ -1127,7 +1173,7 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
if userName == "" {
userName = "user"
}
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port)
forwardID, err := h.repo.CreateForwardTx(userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, inx, entryNodes, port, nullableInt(speedID))
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -1202,6 +1248,24 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
if strategy == "" {
strategy = forward.Strategy
}
speedID := asAnyToInt64Ptr(req["speedId"])
if speedID != nil {
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
if speedErr != nil {
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
return
}
if !exists {
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
return
}
}
newSpeedID := forward.SpeedID
if speedID != nil {
newSpeedID = sql.NullInt64{Int64: *speedID, Valid: true}
} else if _, ok := req["speedId"]; ok {
newSpeedID = sql.NullInt64{Valid: false}
}
port := asInt(req["inPort"], 0)
if port <= 0 {
@@ -1225,7 +1289,7 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
}
}
now := time.Now().UnixMilli()
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now); err != nil {
if err := h.repo.UpdateForward(id, name, tunnelID, remoteAddr, strategy, now, newSpeedID); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -3006,6 +3070,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,
time.Now().UnixMilli(),
)
@@ -3027,6 +3092,10 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
h.repo.GetExistingUserTunnel(userID, tunnelID)
speedID := asAnyToInt64Ptr(req["speedId"])
if err := h.validateSpeedLimitReference(speedID); err != nil {
return err
}
reqFlow := asInt64(req["flow"], -1)
reqNum := asInt(req["num"], -1)
reqExpTime := asInt64(req["expTime"], -1)
@@ -3067,7 +3136,24 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
reqStatus = 1
}
return h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus)
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil {
return err
}
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
insertedID, _, _, _, _, _, _, lookupErr := h.repo.GetExistingUserTunnel(userID, tunnelID)
if lookupErr != nil {
return fmt.Errorf("下发失败且回滚失败: %v; 回滚查询错误: %w", syncErr, lookupErr)
}
if rollbackErr := h.repo.DeleteUserTunnel(insertedID); rollbackErr != nil {
return fmt.Errorf("下发失败且回滚失败: %v; 回滚删除错误: %w", syncErr, rollbackErr)
}
return fmt.Errorf("下发失败,已回滚: %w", syncErr)
}
return nil
}
if err != nil {
return err
@@ -3105,25 +3191,61 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
newSpeedID = sql.NullInt64{Valid: false}
}
err = h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus)
if err == nil {
h.syncUserTunnelForwards(userID, tunnelID)
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
return err
}
return err
if syncErr := h.syncUserTunnelForwards(userID, tunnelID); syncErr != nil {
rollbackErr := h.repo.UpdateUserTunnelFields(
existingID,
currentSpeedID,
currentFlow,
int(currentNum),
currentExpTime,
currentFlowReset,
currentStatus,
)
if rollbackErr != nil {
return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr)
}
return fmt.Errorf("下发失败,已回滚: %w", syncErr)
}
return nil
}
func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) {
func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error {
forwards, err := h.listForwardsByTunnel(tunnelID)
if err != nil {
return
return err
}
for i := range forwards {
f := &forwards[i]
if f.UserID == userID {
_ = h.syncForwardServices(f, "UpdateService", true)
if err := h.syncForwardServices(f, "UpdateService", true); err != nil {
return err
}
}
}
return nil
}
func (h *Handler) validateSpeedLimitReference(speedID *int64) error {
if speedID == nil {
return nil
}
exists, err := h.repo.SpeedLimitExists(*speedID)
if err != nil {
return err
}
if !exists {
return errors.New("限速规则不存在")
}
return nil
}
func asAnySlice(v interface{}) []interface{} {