mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
feat: support adaptive traffic units and precise quotas
This commit is contained in:
@@ -4,6 +4,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log"
|
||||
"math"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
@@ -12,12 +13,27 @@ import (
|
||||
)
|
||||
|
||||
const bytesPerGB int64 = 1024 * 1024 * 1024
|
||||
const bytesPerMiB int64 = 1024 * 1024
|
||||
|
||||
func flowLimitBytes(flowGB, flowMiB int64) int64 {
|
||||
if flowMiB > 0 {
|
||||
if flowMiB > math.MaxInt64/bytesPerMiB {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return flowMiB * bytesPerMiB
|
||||
}
|
||||
if flowGB > math.MaxInt64/bytesPerGB {
|
||||
return math.MaxInt64
|
||||
}
|
||||
return flowGB * bytesPerGB
|
||||
}
|
||||
|
||||
type userTunnelPolicy struct {
|
||||
ID int64
|
||||
UserID int64
|
||||
TunnelID int64
|
||||
Flow int64
|
||||
FlowMiB int64
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
ExpTime int64
|
||||
@@ -358,7 +374,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
return errors.New("账号已过期")
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return errors.New("流量已超额,禁止开启转发")
|
||||
@@ -400,7 +416,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
return errors.New("该隧道已过期")
|
||||
}
|
||||
|
||||
utFlowLimit := policy.Flow * bytesPerGB
|
||||
utFlowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
|
||||
utCurrent := policy.InFlow + policy.OutFlow
|
||||
if utCurrent >= utFlowLimit {
|
||||
return errors.New("该隧道流量已超额,禁止开启转发")
|
||||
@@ -425,7 +441,7 @@ func (h *Handler) shouldPauseUser(userID int64, now int64) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := user.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(user.Flow, user.FlowMiB)
|
||||
current := user.InFlow + user.OutFlow
|
||||
if flowLimit < current {
|
||||
return true
|
||||
@@ -441,7 +457,7 @@ func shouldPauseUserTunnel(policy *userTunnelPolicy, now int64) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
flowLimit := policy.Flow * bytesPerGB
|
||||
flowLimit := flowLimitBytes(policy.Flow, policy.FlowMiB)
|
||||
current := policy.InFlow + policy.OutFlow
|
||||
if current >= flowLimit {
|
||||
return true
|
||||
@@ -465,7 +481,7 @@ func (h *Handler) getUserTunnelPolicy(userTunnelID int64) (*userTunnelPolicy, er
|
||||
}
|
||||
return &userTunnelPolicy{
|
||||
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
|
||||
Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
ExpTime: ut.ExpTime, Status: ut.Status, Num: ut.Num,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -719,6 +719,7 @@ func (h *Handler) userTunnelList(w http.ResponseWriter, r *http.Request) {
|
||||
"tunnelName": t.TunnelName,
|
||||
"status": t.Status,
|
||||
"flow": t.Flow,
|
||||
"flowMiB": t.FlowMiB,
|
||||
"num": t.Num,
|
||||
"expTime": t.ExpTime,
|
||||
"flowResetTime": t.FlowResetTime,
|
||||
@@ -1213,6 +1214,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
"tunnelName": t.TunnelName,
|
||||
"tunnelFlow": t.TunnelFlow,
|
||||
"flow": t.Flow,
|
||||
"flowMiB": t.FlowMiB,
|
||||
"inFlow": t.InFlow,
|
||||
"outFlow": t.OutFlow,
|
||||
"num": t.Num,
|
||||
@@ -1262,6 +1264,7 @@ func (h *Handler) userPackage(w http.ResponseWriter, r *http.Request) {
|
||||
"user": user.User,
|
||||
"status": user.Status,
|
||||
"flow": user.Flow,
|
||||
"flowMiB": user.FlowMiB,
|
||||
"inFlow": user.InFlow,
|
||||
"outFlow": user.OutFlow,
|
||||
"num": user.Num,
|
||||
|
||||
@@ -57,7 +57,11 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
status := asInt(req["status"], 1)
|
||||
flow := asInt64(req["flow"], 100)
|
||||
flow, flowMiB, flowErr := parseTrafficLimit(req, 100)
|
||||
if flowErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
|
||||
return
|
||||
}
|
||||
num := asInt(req["num"], 10)
|
||||
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
|
||||
flowResetTime := asInt64(req["flowResetTime"], 1)
|
||||
@@ -76,7 +80,7 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now)
|
||||
userID, err := h.repo.CreateUser(username, hashedPassword, roleID, expTime, flow, flowResetTime, num, status, maxConn, now, flowMiB)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -164,7 +168,16 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
flow := asInt64(req["flow"], 100)
|
||||
flow, flowMiB, flowErr := parseTrafficLimit(req, 100)
|
||||
if flowErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
|
||||
return
|
||||
}
|
||||
if _, supplied := req["flowMiB"]; !supplied {
|
||||
if current, err := h.repo.GetUserByID(id); err == nil && current != nil && current.Flow == flow {
|
||||
flowMiB = current.FlowMiB
|
||||
}
|
||||
}
|
||||
num := asInt(req["num"], 10)
|
||||
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
|
||||
flowResetTime := asInt64(req["flowResetTime"], 1)
|
||||
@@ -176,7 +189,7 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
pwd := asString(req["pwd"])
|
||||
if strings.TrimSpace(pwd) == "" {
|
||||
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
if err := h.repo.UpdateUserWithoutPassword(id, username, flow, num, expTime, flowResetTime, status, maxConn, now, flowMiB); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -186,13 +199,13 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now); err != nil {
|
||||
if err := h.repo.UpdateUserWithPassword(id, username, hashedPassword, flow, num, expTime, flowResetTime, status, maxConn, now, flowMiB); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime)
|
||||
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime, flowMiB)
|
||||
if hasDailyQuota || hasMonthlyQuota {
|
||||
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
|
||||
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
|
||||
@@ -1991,14 +2004,32 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, oldErr.Error()))
|
||||
return
|
||||
}
|
||||
oldTunnel, oldTunnelErr := h.repo.GetUserTunnelByID(id)
|
||||
if oldTunnelErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, oldTunnelErr.Error()))
|
||||
return
|
||||
}
|
||||
if oldTunnel == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("隧道权限不存在"))
|
||||
return
|
||||
}
|
||||
flow, flowMiB, flowErr := parseTrafficLimit(req, 0)
|
||||
if flowErr != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(flowErr.Error()))
|
||||
return
|
||||
}
|
||||
if _, supplied := req["flowMiB"]; !supplied && oldTunnel.Flow == flow {
|
||||
flowMiB = oldTunnel.FlowMiB
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserTunnel(id,
|
||||
asInt64(req["flow"], 0),
|
||||
flow,
|
||||
asInt(req["num"], 0),
|
||||
asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli()),
|
||||
asInt64(req["flowResetTime"], 1),
|
||||
nullableInt(speedID),
|
||||
asInt(req["status"], 1),
|
||||
flowMiB,
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -2013,6 +2044,7 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
oldFlowReset,
|
||||
oldSpeedID,
|
||||
oldStatus,
|
||||
oldTunnel.FlowMiB,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("下发失败且回滚失败: %v; 回滚错误: %v", syncErr, rollbackErr)))
|
||||
@@ -4781,6 +4813,14 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
}
|
||||
|
||||
reqFlow := asInt64(req["flow"], -1)
|
||||
var reqFlowMiB int64
|
||||
if _, hasFlowMiB := req["flowMiB"]; hasFlowMiB {
|
||||
var flowErr error
|
||||
reqFlow, reqFlowMiB, flowErr = parseTrafficLimit(req, 0)
|
||||
if flowErr != nil {
|
||||
return flowErr
|
||||
}
|
||||
}
|
||||
reqNum := asInt(req["num"], -1)
|
||||
reqExpTime := asInt64(req["expTime"], -1)
|
||||
reqFlowReset := asInt64(req["flowResetTime"], -1)
|
||||
@@ -4792,6 +4832,9 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
if uErr == nil {
|
||||
if reqFlow < 0 {
|
||||
reqFlow = uFlow
|
||||
if user, err := h.repo.GetUserByID(userID); err == nil && user != nil {
|
||||
reqFlowMiB = user.FlowMiB
|
||||
}
|
||||
}
|
||||
if reqNum < 0 {
|
||||
reqNum = uNum
|
||||
@@ -4820,7 +4863,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
reqStatus = 1
|
||||
}
|
||||
|
||||
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus); err != nil {
|
||||
if err := h.repo.InsertUserTunnel(userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus, reqFlowMiB); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -4844,8 +4887,20 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
}
|
||||
|
||||
newFlow := currentFlow
|
||||
oldTunnel, err := h.repo.GetUserTunnelByID(existingID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if oldTunnel == nil {
|
||||
return fmt.Errorf("隧道权限不存在")
|
||||
}
|
||||
newFlowMiB := oldTunnel.FlowMiB
|
||||
if reqFlow >= 0 {
|
||||
newFlow = reqFlow
|
||||
newFlowMiB = reqFlowMiB
|
||||
if _, supplied := req["flowMiB"]; !supplied && reqFlow == currentFlow {
|
||||
newFlowMiB = oldTunnel.FlowMiB
|
||||
}
|
||||
}
|
||||
|
||||
newNum := int(currentNum)
|
||||
@@ -4875,7 +4930,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus); err != nil {
|
||||
if err := h.repo.UpdateUserTunnelFields(existingID, newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus, newFlowMiB); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -4888,6 +4943,7 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
currentExpTime,
|
||||
currentFlowReset,
|
||||
currentStatus,
|
||||
oldTunnel.FlowMiB,
|
||||
)
|
||||
if rollbackErr != nil {
|
||||
return fmt.Errorf("下发失败且回滚失败: %v; 回滚错误: %w", syncErr, rollbackErr)
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// flowMiB is optional so older clients can keep sending the GB-based flow field.
|
||||
// A positive value takes precedence and preserves sub-GB limits exactly.
|
||||
func parseTrafficLimit(req map[string]interface{}, defaultGB int64) (flowGB, flowMiB int64, err error) {
|
||||
flowGB = asInt64(req["flow"], defaultGB)
|
||||
if flowGB < 0 {
|
||||
return 0, 0, fmt.Errorf("流量限制不能小于0")
|
||||
}
|
||||
raw, present := req["flowMiB"]
|
||||
if !present {
|
||||
return flowGB, 0, nil
|
||||
}
|
||||
flowMiB, err = strconv.ParseInt(asString(raw), 10, 64)
|
||||
if err != nil || flowMiB < 0 || flowMiB > math.MaxInt64/bytesPerMiB {
|
||||
return 0, 0, fmt.Errorf("流量限制超出范围")
|
||||
}
|
||||
if flowMiB > 0 {
|
||||
flowGB = (flowMiB-1)/1024 + 1
|
||||
}
|
||||
return flowGB, flowMiB, nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestTrafficLimitMiBOverridesLegacyGB(t *testing.T) {
|
||||
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{
|
||||
"flow": float64(1), "flowMiB": float64(500),
|
||||
}, 100)
|
||||
if err != nil || flowGB != 1 || flowMiB != 500 {
|
||||
t.Fatalf("parseTrafficLimit = (%d, %d, %v), want (1, 500, nil)", flowGB, flowMiB, err)
|
||||
}
|
||||
limit := flowLimitBytes(flowGB, flowMiB)
|
||||
if limit != 500*bytesPerMiB {
|
||||
t.Fatalf("limit = %d, want %d", limit, 500*bytesPerMiB)
|
||||
}
|
||||
policy := &userTunnelPolicy{Flow: flowGB, FlowMiB: flowMiB, InFlow: limit - 1, Status: 1}
|
||||
if shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
|
||||
t.Fatal("policy paused before reaching 500 MiB")
|
||||
}
|
||||
policy.InFlow = limit
|
||||
if !shouldPauseUserTunnel(policy, time.Now().UnixMilli()) {
|
||||
t.Fatal("policy did not pause at 500 MiB")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrafficLimitLegacyAndInvalidValues(t *testing.T) {
|
||||
flowGB, flowMiB, err := parseTrafficLimit(map[string]interface{}{"flow": float64(2)}, 100)
|
||||
if err != nil || flowGB != 2 || flowMiB != 0 || flowLimitBytes(flowGB, flowMiB) != 2*bytesPerGB {
|
||||
t.Fatalf("legacy GB limit changed: (%d, %d, %v)", flowGB, flowMiB, err)
|
||||
}
|
||||
for _, value := range []interface{}{"1.5", -1, "999999999999999999999"} {
|
||||
if _, _, err := parseTrafficLimit(map[string]interface{}{"flowMiB": value}, 100); err == nil {
|
||||
t.Fatalf("accepted invalid flowMiB %v", value)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user