mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-06 10:06:36 +08:00
feat: add speed limit contract tests and refine limit/user UI
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{} {
|
||||
|
||||
Reference in New Issue
Block a user