mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-06 10:06:36 +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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -16,6 +16,7 @@ type User struct {
|
||||
RoleID int `gorm:"column:role_id;not null"`
|
||||
ExpTime int64 `gorm:"column:exp_time;not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
FlowMiB int64 `gorm:"column:flow_mib;not null;default:0"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
@@ -230,6 +231,7 @@ type UserTunnel struct {
|
||||
SpeedID sql.NullInt64 `gorm:"column:speed_id"`
|
||||
Num int `gorm:"not null"`
|
||||
Flow int64 `gorm:"not null"`
|
||||
FlowMiB int64 `gorm:"column:flow_mib;not null;default:0"`
|
||||
InFlow int64 `gorm:"column:in_flow;not null;default:0"`
|
||||
OutFlow int64 `gorm:"column:out_flow;not null;default:0"`
|
||||
FlowResetTime int64 `gorm:"column:flow_reset_time;not null"`
|
||||
@@ -417,6 +419,7 @@ type UserBackup struct {
|
||||
RoleID int `json:"roleId"`
|
||||
ExpTime int64 `json:"expTime"`
|
||||
Flow int64 `json:"flow"`
|
||||
FlowMiB int64 `json:"flowMiB,omitempty"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
@@ -523,6 +526,7 @@ type UserTunnelBackup struct {
|
||||
SpeedID int64 `json:"speedId,omitempty"`
|
||||
Num int `json:"num"`
|
||||
Flow int64 `json:"flow"`
|
||||
FlowMiB int64 `json:"flowMiB,omitempty"`
|
||||
InFlow int64 `json:"inFlow"`
|
||||
OutFlow int64 `json:"outFlow"`
|
||||
FlowResetTime int64 `json:"flowResetTime"`
|
||||
@@ -706,6 +710,7 @@ type UserTunnelDetail struct {
|
||||
Status int
|
||||
TunnelFlow int
|
||||
Flow int64
|
||||
FlowMiB int64 `gorm:"column:flow_mib"`
|
||||
InFlow int64
|
||||
OutFlow int64
|
||||
Num int
|
||||
|
||||
@@ -669,7 +669,7 @@ func (r *Repository) GetUserPackageTunnels(userID int64) ([]model.UserTunnelDeta
|
||||
}
|
||||
var items []model.UserTunnelDetail
|
||||
err := r.db.Model(&model.UserTunnel{}).
|
||||
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
|
||||
Select("user_tunnel.id, user_tunnel.user_id, user_tunnel.tunnel_id, tunnel.name AS tunnel_name, user_tunnel.status, tunnel.flow AS tunnel_flow, user_tunnel.flow, user_tunnel.flow_mib, user_tunnel.in_flow, user_tunnel.out_flow, user_tunnel.num, user_tunnel.flow_reset_time, user_tunnel.exp_time, user_tunnel.speed_id, speed_limit.name AS speed_limit, speed_limit.speed").
|
||||
Joins("LEFT JOIN tunnel ON tunnel.id = user_tunnel.tunnel_id").
|
||||
Joins("LEFT JOIN speed_limit ON speed_limit.id = user_tunnel.speed_id").
|
||||
Where("user_tunnel.user_id = ?", userID).
|
||||
@@ -924,7 +924,7 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
|
||||
item := map[string]interface{}{
|
||||
"id": u.ID, "user": u.User, "name": u.User,
|
||||
"roleId": u.RoleID, "status": u.Status,
|
||||
"flow": u.Flow, "num": u.Num, "expTime": u.ExpTime,
|
||||
"flow": u.Flow, "flowMiB": u.FlowMiB, "num": u.Num, "expTime": u.ExpTime,
|
||||
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
|
||||
"updatedTime": nullableInt64(u.UpdatedTime),
|
||||
"inFlow": u.InFlow, "outFlow": u.OutFlow,
|
||||
@@ -2094,7 +2094,7 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
|
||||
for _, u := range users {
|
||||
b := model.UserBackup{
|
||||
ID: u.ID, User: u.User, Pwd: u.Pwd, RoleID: u.RoleID,
|
||||
ExpTime: u.ExpTime, Flow: u.Flow, InFlow: u.InFlow, OutFlow: u.OutFlow,
|
||||
ExpTime: u.ExpTime, Flow: u.Flow, FlowMiB: u.FlowMiB, InFlow: u.InFlow, OutFlow: u.OutFlow,
|
||||
FlowResetTime: u.FlowResetTime, Num: u.Num,
|
||||
CreatedTime: u.CreatedTime, Status: u.Status,
|
||||
}
|
||||
@@ -2270,7 +2270,7 @@ func (r *Repository) exportUserTunnels() ([]model.UserTunnelBackup, error) {
|
||||
for _, ut := range uts {
|
||||
b := model.UserTunnelBackup{
|
||||
ID: ut.ID, UserID: ut.UserID, TunnelID: ut.TunnelID,
|
||||
Num: ut.Num, Flow: ut.Flow, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
Num: ut.Num, Flow: ut.Flow, FlowMiB: ut.FlowMiB, InFlow: ut.InFlow, OutFlow: ut.OutFlow,
|
||||
FlowResetTime: ut.FlowResetTime, ExpTime: ut.ExpTime, Status: ut.Status,
|
||||
}
|
||||
if ut.SpeedID.Valid {
|
||||
@@ -2467,6 +2467,7 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
|
||||
RoleID: u.RoleID,
|
||||
ExpTime: u.ExpTime,
|
||||
Flow: u.Flow,
|
||||
FlowMiB: u.FlowMiB,
|
||||
InFlow: u.InFlow,
|
||||
OutFlow: u.OutFlow,
|
||||
FlowResetTime: u.FlowResetTime,
|
||||
@@ -2479,7 +2480,7 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
|
||||
err = tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user", "pwd", "role_id", "exp_time", "flow", "in_flow", "out_flow",
|
||||
"user", "pwd", "role_id", "exp_time", "flow", "flow_mib", "in_flow", "out_flow",
|
||||
"flow_reset_time", "num", "updated_time", "status", "password_changed_at",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
@@ -2717,6 +2718,7 @@ func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int6
|
||||
SpeedID: sql.NullInt64{Int64: ut.SpeedID, Valid: ut.SpeedID > 0},
|
||||
Num: ut.Num,
|
||||
Flow: ut.Flow,
|
||||
FlowMiB: ut.FlowMiB,
|
||||
InFlow: ut.InFlow,
|
||||
OutFlow: ut.OutFlow,
|
||||
FlowResetTime: ut.FlowResetTime,
|
||||
@@ -2726,7 +2728,7 @@ func importUserTunnels(tx *gorm.DB, userTunnels []model.UserTunnelBackup, _ int6
|
||||
err := tx.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"user_id", "tunnel_id", "speed_id", "num", "flow", "in_flow", "out_flow",
|
||||
"user_id", "tunnel_id", "speed_id", "num", "flow", "flow_mib", "in_flow", "out_flow",
|
||||
"flow_reset_time", "exp_time", "status",
|
||||
}),
|
||||
}).Create(&item).Error
|
||||
|
||||
@@ -37,7 +37,14 @@ func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool
|
||||
return cnt > 0, err
|
||||
}
|
||||
|
||||
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64) (int64, error) {
|
||||
func optionalFlowMiB(values []int64) int64 {
|
||||
if len(values) > 0 {
|
||||
return values[0]
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status, maxConn int, now int64, flowMiB ...int64) (int64, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return 0, errors.New("repository not initialized")
|
||||
}
|
||||
@@ -47,6 +54,7 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
|
||||
RoleID: roleID,
|
||||
ExpTime: expTime,
|
||||
Flow: flow,
|
||||
FlowMiB: optionalFlowMiB(flowMiB),
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
@@ -75,7 +83,7 @@ func (r *Repository) GetUserRoleID(userID int64) (int, error) {
|
||||
return user.RoleID, nil
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
|
||||
func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -85,6 +93,7 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
|
||||
"user": username,
|
||||
"pwd": pwdHash,
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -95,7 +104,7 @@ func (r *Repository) UpdateUserWithPassword(id int64, username, pwdHash string,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64) error {
|
||||
func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow int64, num int, expTime, flowResetTime int64, status, maxConn int, now int64, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -104,6 +113,7 @@ func (r *Repository) UpdateUserWithoutPassword(id int64, username string, flow i
|
||||
Updates(map[string]interface{}{
|
||||
"user": username,
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -126,7 +136,7 @@ func (r *Repository) UpdateUserPassword(userID int64, pwdHash string, now int64)
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64) {
|
||||
func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num int, expTime, flowResetTime int64, flowMiB ...int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
}
|
||||
@@ -134,6 +144,7 @@ func (r *Repository) PropagateUserFlowToTunnels(userID int64, flow int64, num in
|
||||
Where("user_id = ?", userID).
|
||||
Updates(map[string]interface{}{
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -650,7 +661,7 @@ func (r *Repository) DeleteUserTunnel(id int64) error {
|
||||
return r.db.Where("id = ?", id).Delete(&model.UserTunnel{}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int) error {
|
||||
func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, flowResetTime int64, speedID interface{}, status int, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -658,6 +669,7 @@ func (r *Repository) UpdateUserTunnel(id int64, flow int64, num int, expTime, fl
|
||||
Where("id = ?", id).
|
||||
Updates(map[string]interface{}{
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -692,7 +704,7 @@ func (r *Repository) GetExistingUserTunnel(userID, tunnelID int64) (id int64, fl
|
||||
return ut.ID, ut.Flow, int64(ut.Num), ut.ExpTime, ut.FlowResetTime, ut.SpeedID, ut.Status, nil
|
||||
}
|
||||
|
||||
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int) error {
|
||||
func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{}, num int, flow, flowResetTime, expTime int64, status int, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -702,6 +714,7 @@ func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{
|
||||
SpeedID: nullInt64FromInterface(speedID),
|
||||
Num: num,
|
||||
Flow: flow,
|
||||
FlowMiB: optionalFlowMiB(flowMiB),
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowResetTime,
|
||||
@@ -711,7 +724,7 @@ func (r *Repository) InsertUserTunnel(userID, tunnelID int64, speedID interface{
|
||||
return r.db.Create(&ut).Error
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int) error {
|
||||
func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow int64, num int, expTime, flowResetTime int64, status int, flowMiB ...int64) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
@@ -720,6 +733,7 @@ func (r *Repository) UpdateUserTunnelFields(id int64, speedID interface{}, flow
|
||||
Updates(map[string]interface{}{
|
||||
"speed_id": nullInt64FromInterface(speedID),
|
||||
"flow": flow,
|
||||
"flow_mib": optionalFlowMiB(flowMiB),
|
||||
"num": num,
|
||||
"exp_time": expTime,
|
||||
"flow_reset_time": flowResetTime,
|
||||
@@ -1293,7 +1307,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
return 0, false, err
|
||||
}
|
||||
var user model.User
|
||||
if err := r.db.Select("flow, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
if err := r.db.Select("flow, flow_mib, num, exp_time, flow_reset_time").Where("id = ?", userID).First(&user).Error; err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
flow := user.Flow
|
||||
@@ -1305,6 +1319,7 @@ func (r *Repository) EnsureUserTunnelGrant(userID, tunnelID int64) (int64, bool,
|
||||
TunnelID: tunnelID,
|
||||
Num: num,
|
||||
Flow: flow,
|
||||
FlowMiB: user.FlowMiB,
|
||||
InFlow: 0,
|
||||
OutFlow: 0,
|
||||
FlowResetTime: flowReset,
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
)
|
||||
|
||||
func TestTrafficLimitMiBSurvivesBackupRestore(t *testing.T) {
|
||||
source, err := Open(filepath.Join(t.TempDir(), "source.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer source.Close()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
userID, err := source.CreateUser("mib-user", "hash", 1, now+86400000, 1, 1, 10, 1, 0, now, 500)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tunnel := model.Tunnel{Name: "mib-tunnel", TrafficRatio: 1, Type: 1, Protocol: "tls", Flow: 1, CreatedTime: now, UpdatedTime: now, Status: 1, Inx: 1}
|
||||
if err := source.DB().Create(&tunnel).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, _, err := source.EnsureUserTunnelGrant(userID, tunnel.ID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
grants, err := source.GetUserPackageTunnels(userID)
|
||||
if err != nil || len(grants) != 1 || grants[0].FlowMiB != 500 {
|
||||
t.Fatalf("inherited tunnel quota = %+v, err = %v", grants, err)
|
||||
}
|
||||
backup, err := source.ExportAll()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
found := false
|
||||
for _, user := range backup.Users {
|
||||
if user.User == "mib-user" {
|
||||
found = user.Flow == 1 && user.FlowMiB == 500
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("500 MiB user quota missing from backup")
|
||||
}
|
||||
if len(backup.UserTunnels) != 1 || backup.UserTunnels[0].FlowMiB != 500 {
|
||||
t.Fatalf("tunnel quota missing from backup: %+v", backup.UserTunnels)
|
||||
}
|
||||
|
||||
dest, err := Open(filepath.Join(t.TempDir(), "dest.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer dest.Close()
|
||||
if _, err := dest.Import(backup, []string{"users", "tunnels", "userTunnels"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
user, err := dest.GetUserByUsername("mib-user")
|
||||
if err != nil || user == nil || user.Flow != 1 || user.FlowMiB != 500 {
|
||||
t.Fatalf("restored quota = %+v, err = %v", user, err)
|
||||
}
|
||||
grants, err = dest.GetUserPackageTunnels(user.ID)
|
||||
if err != nil || len(grants) != 1 || grants[0].FlowMiB != 500 {
|
||||
t.Fatalf("restored tunnel quota = %+v, err = %v", grants, err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user