diff --git a/go-backend/internal/http/handler/flow_policy.go b/go-backend/internal/http/handler/flow_policy.go index cb3f472..cca261f 100644 --- a/go-backend/internal/http/handler/flow_policy.go +++ b/go-backend/internal/http/handler/flow_policy.go @@ -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 } diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index 1c87fb8..4d67883 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -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, diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 8b9f4b9..1a11ab1 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -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) diff --git a/go-backend/internal/http/handler/traffic_limit.go b/go-backend/internal/http/handler/traffic_limit.go new file mode 100644 index 0000000..d23b587 --- /dev/null +++ b/go-backend/internal/http/handler/traffic_limit.go @@ -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 +} diff --git a/go-backend/internal/http/handler/traffic_limit_test.go b/go-backend/internal/http/handler/traffic_limit_test.go new file mode 100644 index 0000000..f223422 --- /dev/null +++ b/go-backend/internal/http/handler/traffic_limit_test.go @@ -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) + } + } +} diff --git a/go-backend/internal/store/model/model.go b/go-backend/internal/store/model/model.go index d35cc8a..afca82b 100644 --- a/go-backend/internal/store/model/model.go +++ b/go-backend/internal/store/model/model.go @@ -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 diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index db68830..a238d63 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -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 diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index 5b3685d..34e0bdb 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -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, diff --git a/go-backend/internal/store/repo/repository_traffic_limit_test.go b/go-backend/internal/store/repo/repository_traffic_limit_test.go new file mode 100644 index 0000000..a3ce44e --- /dev/null +++ b/go-backend/internal/store/repo/repository_traffic_limit_test.go @@ -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) + } +} diff --git a/vite-frontend/src/api/types.ts b/vite-frontend/src/api/types.ts index 5cdc3b9..a33cb5e 100644 --- a/vite-frontend/src/api/types.ts +++ b/vite-frontend/src/api/types.ts @@ -19,6 +19,7 @@ export interface UserApiItem { name?: string; status: number; flow: number; + flowMiB?: number; num: number; expTime?: number; flowResetTime?: number; @@ -106,6 +107,7 @@ export interface UserTunnelPermissionApiItem { tunnelName: string; status: number; flow: number; + flowMiB?: number; num: number; expTime: number; flowResetTime: number; @@ -387,6 +389,7 @@ export interface UserTunnelAssignPayload { id?: number; tunnelId?: number; flow?: number; + flowMiB?: number; num?: number; expTime?: number; flowResetTime?: number; diff --git a/vite-frontend/src/components/traffic-limit-field.tsx b/vite-frontend/src/components/traffic-limit-field.tsx new file mode 100644 index 0000000..01a7917 --- /dev/null +++ b/vite-frontend/src/components/traffic-limit-field.tsx @@ -0,0 +1,62 @@ +import { Input } from "@/shadcn-bridge/heroui/input"; +import { Select, SelectItem } from "@/shadcn-bridge/heroui/select"; +import { + parseTrafficInput, + TRAFFIC_UNIT_MIB, + type TrafficUnit, +} from "@/utils/traffic"; + +const UNITS: TrafficUnit[] = ["MB", "GB", "TB", "PB"]; + +interface TrafficLimitFieldProps { + label: string; + value: string; + unit: TrafficUnit; + onChange: (value: string, unit: TrafficUnit) => void; + description?: string; + isRequired?: boolean; +} + +export function TrafficLimitField({ + label, + value, + unit, + onChange, + description, + isRequired, +}: TrafficLimitFieldProps) { + return ( +
+ onChange(event.target.value, unit)} + /> + +
+ ); +} diff --git a/vite-frontend/src/pages/dashboard.tsx b/vite-frontend/src/pages/dashboard.tsx index 5b63697..cc1a810 100644 --- a/vite-frontend/src/pages/dashboard.tsx +++ b/vite-frontend/src/pages/dashboard.tsx @@ -24,6 +24,11 @@ import { FlowChartCard } from "@/pages/dashboard/components/flow-chart-card"; import { MetricCard } from "@/pages/dashboard/components/metric-card"; import { getSessionName } from "@/utils/session"; import { safeLogout } from "@/utils/logout"; +import { + formatTraffic, + formatFlowLimit, + flowLimitBytes, +} from "@/utils/traffic"; import { formatNodeRenewalTime, getNodeRenewalCycleLabel, @@ -70,24 +75,7 @@ export default function DashboardPage() { const [addressModalTitle, setAddressModalTitle] = useState(""); const [addressList, setAddressList] = useState([]); - const formatFlow = (value: number, unit: string = "bytes"): string => { - // 99999 表示无限制 - if (value === 99999) { - return "无限制"; - } - - if (unit === "gb") { - return value + " GB"; - } else { - if (value === 0) return "0 B"; - if (value < 1024) return value + " B"; - if (value < 1024 * 1024) return (value / 1024).toFixed(2) + " KB"; - if (value < 1024 * 1024 * 1024) - return (value / (1024 * 1024)).toFixed(2) + " MB"; - - return (value / (1024 * 1024 * 1024)).toFixed(2) + " GB"; - } - }; + const formatFlow = formatTraffic; const formatNumber = (value: number): string => { // 99999 表示无限制 @@ -279,10 +267,10 @@ export default function DashboardPage() { const calculateUsagePercentage = (type: "flow" | "forwards"): number => { if (type === "flow") { const totalUsed = calculateUserTotalUsedFlow(); - const totalLimit = (userInfo.flow || 0) * 1024 * 1024 * 1024; + const totalLimit = flowLimitBytes(userInfo.flow || 0, userInfo.flowMiB); // 无限制时返回0% - if (userInfo.flow === 99999) return 0; + if (userInfo.flow === 99999 && !userInfo.flowMiB) return 0; return totalLimit > 0 ? Math.min((totalUsed / totalLimit) * 100, 100) : 0; } else if (type === "forwards") { @@ -351,10 +339,10 @@ export default function DashboardPage() { const calculateTunnelFlowPercentage = (tunnel: UserTunnel): number => { const totalUsed = calculateTunnelUsedFlow(tunnel); - const totalLimit = (tunnel.flow || 0) * 1024 * 1024 * 1024; + const totalLimit = flowLimitBytes(tunnel.flow || 0, tunnel.flowMiB); // 无限制时返回0% - if (tunnel.flow === 99999) return 0; + if (tunnel.flow === 99999 && !tunnel.flowMiB) return 0; return totalLimit > 0 ? Math.min((totalUsed / totalLimit) * 100, 100) : 0; }; @@ -741,7 +729,7 @@ export default function DashboardPage() { } iconClassName="bg-blue-100 dark:bg-blue-500/20" title="总流量" - value={formatFlow(userInfo.flow, "gb")} + value={formatFlowLimit(userInfo.flow, userInfo.flowMiB)} />

- {userInfo.flow === 99999 + {userInfo.flow === 99999 && !userInfo.flowMiB ? "无限制" : `${calculateUsagePercentage("flow").toFixed(1)}%`}

@@ -981,7 +969,7 @@ export default function DashboardPage() { 流量配额

- {formatFlow(tunnel.flow, "gb")} + {formatFlowLimit(tunnel.flow, tunnel.flowMiB)}

@@ -995,7 +983,7 @@ export default function DashboardPage() { {renderProgressBar( calculateTunnelFlowPercentage(tunnel), "sm", - tunnel.flow === 99999, + tunnel.flow === 99999 && !tunnel.flowMiB, )}
diff --git a/vite-frontend/src/pages/dashboard/use-dashboard-data.ts b/vite-frontend/src/pages/dashboard/use-dashboard-data.ts index bd2eed4..14053e1 100644 --- a/vite-frontend/src/pages/dashboard/use-dashboard-data.ts +++ b/vite-frontend/src/pages/dashboard/use-dashboard-data.ts @@ -14,6 +14,7 @@ import { getAdminFlag } from "@/utils/session"; export interface DashboardUserInfo { flow: number; + flowMiB?: number; inFlow: number; outFlow: number; num: number; @@ -26,6 +27,7 @@ export interface DashboardUserTunnel { tunnelId: number; tunnelName: string; flow: number; + flowMiB?: number; inFlow: number; outFlow: number; num: number; diff --git a/vite-frontend/src/pages/forward.tsx b/vite-frontend/src/pages/forward.tsx index f7fc3ba..dfb44e4 100644 --- a/vite-frontend/src/pages/forward.tsx +++ b/vite-frontend/src/pages/forward.tsx @@ -27,6 +27,7 @@ import { useSortable } from "@dnd-kit/sortable"; import { CSS } from "@dnd-kit/utilities"; import { AnimatedPage } from "@/components/animated-page"; +import { formatTraffic } from "@/utils/traffic"; import { BatchActionResultModal } from "@/components/batch-action-result-modal"; import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card"; import { Button } from "@/shadcn-bridge/heroui/button"; @@ -2704,13 +2705,7 @@ export default function ForwardPage() { // 格式化流量 const formatFlow = (value: number): string => { - if (value === 0) return "0 B"; - if (value < 1024) return value + " B"; - if (value < 1024 * 1024) return (value / 1024).toFixed(2) + " KB"; - if (value < 1024 * 1024 * 1024) - return (value / (1024 * 1024)).toFixed(2) + " MB"; - - return (value / (1024 * 1024 * 1024)).toFixed(2) + " GB"; + return formatTraffic(value); }; // 显示地址列表弹窗 diff --git a/vite-frontend/src/pages/node.tsx b/vite-frontend/src/pages/node.tsx index 3c4aabd..aca4328 100644 --- a/vite-frontend/src/pages/node.tsx +++ b/vite-frontend/src/pages/node.tsx @@ -20,6 +20,7 @@ import { CSS } from "@dnd-kit/utilities"; import { LayoutGrid, List } from "lucide-react"; import { SearchBar } from "@/components/search-bar"; +import { formatTraffic } from "@/utils/traffic"; import { AnimatedPage } from "@/components/animated-page"; import { Table, @@ -734,15 +735,7 @@ export default function NodePage() { // 格式化流量 const formatFlow = (bytes: number): string => { - if (!Number.isFinite(bytes) || bytes <= 0) { - return "0 B"; - } - if (bytes < 1024) return `${bytes} B`; - if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(2)} KB`; - if (bytes < 1024 * 1024 * 1024) - return `${(bytes / (1024 * 1024)).toFixed(2)} MB`; - - return `${(bytes / (1024 * 1024 * 1024)).toFixed(2)} GB`; + return formatTraffic(bytes); }; const formatChainType = (chainType: number, hopInx: number) => { diff --git a/vite-frontend/src/pages/node/monitor-view.tsx b/vite-frontend/src/pages/node/monitor-view.tsx index f08ecee..2dd10a9 100644 --- a/vite-frontend/src/pages/node/monitor-view.tsx +++ b/vite-frontend/src/pages/node/monitor-view.tsx @@ -33,6 +33,7 @@ import { } from "lucide-react"; import toast from "react-hot-toast"; +import { formatTraffic } from "@/utils/traffic"; import { DistroIcon, parseDistroFromVersion, @@ -137,15 +138,7 @@ const formatDateTime = (ts: number): string => { }); }; -const formatBytes = (bytes: number): string => { - if (!Number.isFinite(bytes) || bytes <= 0) return "0 B"; - - const k = 1024; - const sizes = ["B", "KB", "MB", "GB", "TB"]; - const i = Math.floor(Math.log(bytes) / Math.log(k)); - - return `${parseFloat((bytes / Math.pow(k, i)).toFixed(2))} ${sizes[i]}`; -}; +const formatBytes = formatTraffic; const formatBytesPerSecond = (bytesPerSecond: number): string => { if (!Number.isFinite(bytesPerSecond) || bytesPerSecond <= 0) return "0 B/s"; diff --git a/vite-frontend/src/pages/node/tunnel-monitor-view.tsx b/vite-frontend/src/pages/node/tunnel-monitor-view.tsx index 5ef5da0..6107513 100644 --- a/vite-frontend/src/pages/node/tunnel-monitor-view.tsx +++ b/vite-frontend/src/pages/node/tunnel-monitor-view.tsx @@ -36,6 +36,7 @@ import { } from "lucide-react"; import toast from "react-hot-toast"; +import { formatTraffic } from "@/utils/traffic"; import { getMonitorTunnels, getTunnelMetrics, @@ -390,12 +391,7 @@ const TrafficChartCard = React.memo(function TrafficChartCard({ const yFormatter = (value: unknown) => { const n = Number(value); - if (!Number.isFinite(n) || n <= 0) return "0 B"; - const k = 1024; - const sizes = ["B", "KB", "MB", "GB", "TB"]; - const i = Math.floor(Math.log(n) / Math.log(k)); - - return `${parseFloat((n / Math.pow(k, i)).toFixed(2))} ${sizes[i]}`; + return formatTraffic(n); }; return ( diff --git a/vite-frontend/src/pages/panel-sharing.tsx b/vite-frontend/src/pages/panel-sharing.tsx index 85b4a81..ee776a3 100644 --- a/vite-frontend/src/pages/panel-sharing.tsx +++ b/vite-frontend/src/pages/panel-sharing.tsx @@ -5,6 +5,15 @@ import { Button } from "@/shadcn-bridge/heroui/button"; import { Card, CardBody, CardHeader } from "@/shadcn-bridge/heroui/card"; import { Tabs, Tab } from "@/shadcn-bridge/heroui/tabs"; import { Input } from "@/shadcn-bridge/heroui/input"; +import { TrafficLimitField } from "@/components/traffic-limit-field"; +import { + formatTraffic, + MIB, + parseTrafficInput, + preferredTrafficUnit, + TRAFFIC_UNIT_MIB, + type TrafficUnit, +} from "@/utils/traffic"; import { Modal, ModalContent, @@ -82,6 +91,8 @@ interface RemoteUsageNode { syncError?: string; } +const MAX_SAFE_BANDWIDTH_MIB = Math.floor(Number.MAX_SAFE_INTEGER / MIB); + export default function PanelSharingPage() { const [selectedTab, setSelectedTab] = useState("my-shares"); const [shares, setShares] = useState([]); @@ -108,6 +119,7 @@ export default function PanelSharingPage() { allowedDomains: "", allowedIps: "", }); + const [shareUnit, setShareUnit] = useState("GB"); const [importForm, setImportForm] = useState({ remoteUrl: "", @@ -124,6 +136,9 @@ export default function PanelSharingPage() { allowedDomains: "", allowedIps: "", }); + const [editUnit, setEditUnit] = useState("GB"); + const [editOriginalMaxBandwidth, setEditOriginalMaxBandwidth] = useState(0); + const [editBandwidthChanged, setEditBandwidthChanged] = useState(false); const loadShares = useCallback(async () => { setLoading(true); @@ -212,8 +227,17 @@ export default function PanelSharingPage() { return; } - if (shareForm.maxBandwidth < 0) { - toast.error("流量上限不能为负数"); + const limitMiB = + shareForm.maxBandwidth === 0 + ? 0 + : parseTrafficInput(String(shareForm.maxBandwidth), shareUnit); + + if ( + limitMiB === null || + limitMiB > MAX_SAFE_BANDWIDTH_MIB || + shareForm.maxBandwidth < 0 + ) { + toast.error("请输入有效的流量上限,0 表示不限流量"); return; } @@ -223,7 +247,7 @@ export default function PanelSharingPage() { const res = await createPeerShare({ name: shareForm.name, nodeId, - maxBandwidth: Math.max(0, shareForm.maxBandwidth) * 1024 * 1024 * 1024, + maxBandwidth: limitMiB * MIB, expiryTime: shareForm.expiryDays === 0 ? 0 : expiryTime, portRangeStart: shareForm.portRangeStart, portRangeEnd: shareForm.portRangeEnd, @@ -274,13 +298,16 @@ export default function PanelSharingPage() { }; const openEditShare = (share: PeerShare) => { + const mib = share.maxBandwidth / MIB; + const unit = preferredTrafficUnit(mib); + + setEditUnit(unit); + setEditOriginalMaxBandwidth(share.maxBandwidth); + setEditBandwidthChanged(false); setEditForm({ id: share.id, name: share.name, - maxBandwidth: - share.maxBandwidth > 0 - ? Math.round(share.maxBandwidth / (1024 * 1024 * 1024)) - : 0, + maxBandwidth: share.maxBandwidth > 0 ? mib / TRAFFIC_UNIT_MIB[unit] : 0, expiryTime: share.expiryTime, portRangeStart: share.portRangeStart, portRangeEnd: share.portRangeEnd, @@ -296,8 +323,18 @@ export default function PanelSharingPage() { return; } - if (editForm.maxBandwidth < 0) { - toast.error("流量上限不能为负数"); + const limitMiB = + editForm.maxBandwidth === 0 + ? 0 + : parseTrafficInput(String(editForm.maxBandwidth), editUnit); + + if ( + editBandwidthChanged && + (limitMiB === null || + limitMiB > MAX_SAFE_BANDWIDTH_MIB || + editForm.maxBandwidth < 0) + ) { + toast.error("请输入有效的流量上限,0 表示不限流量"); return; } @@ -305,7 +342,9 @@ export default function PanelSharingPage() { const res = await updatePeerShare({ id: editForm.id, name: editForm.name, - maxBandwidth: Math.max(0, editForm.maxBandwidth) * 1024 * 1024 * 1024, + maxBandwidth: editBandwidthChanged + ? (limitMiB as number) * MIB + : editOriginalMaxBandwidth, expiryTime: editForm.expiryTime, portRangeStart: editForm.portRangeStart, portRangeEnd: editForm.portRangeEnd, @@ -362,17 +401,7 @@ export default function PanelSharingPage() { toast.success("Token已复制"); }; - const formatFlowGB = (bytes: number) => { - if (!Number.isFinite(bytes) || bytes <= 0) { - return "0 B"; - } - if (bytes < 1024) return bytes + " B"; - if (bytes < 1024 * 1024) return (bytes / 1024).toFixed(2) + " KB"; - if (bytes < 1024 * 1024 * 1024) - return (bytes / (1024 * 1024)).toFixed(2) + " MB"; - - return (bytes / (1024 * 1024 * 1024)).toFixed(2) + " GB"; - }; + const formatFlowGB = formatTraffic; const formatChainType = (chainType: number, hopInx: number) => { if (chainType === 1) { @@ -741,17 +770,18 @@ export default function PanelSharingPage() { }) } /> - - setShareForm({ - ...shareForm, - maxBandwidth: parseInt(e.target.value, 10) || 0, - }) - } + onChange={(value, unit) => { + setShareForm((prev) => ({ + ...prev, + maxBandwidth: Number(value) || 0, + })); + setShareUnit(unit); + }} /> - - setEditForm({ - ...editForm, - maxBandwidth: parseInt(e.target.value, 10) || 0, - }) - } + onChange={(value, unit) => { + setEditForm((prev) => ({ + ...prev, + maxBandwidth: Number(value) || 0, + })); + setEditUnit(unit); + setEditBandwidthChanged(true); + }} /> { - if (unit === "gb") { - return `${value} GB`; - } else { - if (value === 0) return "0 B"; - if (value < 1024) return `${value} B`; - if (value < 1024 * 1024) return `${(value / 1024).toFixed(2)} KB`; - if (value < 1024 * 1024 * 1024) - return `${(value / (1024 * 1024)).toFixed(2)} MB`; - - return `${(value / (1024 * 1024 * 1024)).toFixed(2)} GB`; - } -}; +const formatFlow = formatTraffic; const formatQuotaLimit = (value?: number): string => { const limit = Number(value ?? 0); @@ -96,7 +95,14 @@ const formatQuotaLimit = (value?: number): string => { return "不限"; } - return `${limit} GB`; + return formatTraffic(limit * 1024 ** 3); +}; + +const trafficInputFor = (flowGB: number, flowMiB?: number) => { + const mib = flowLimitMiB(flowGB, flowMiB); + const unit = preferredTrafficUnit(mib); + + return { value: String(mib / TRAFFIC_UNIT_MIB[unit]), unit }; }; const formatDate = (timestamp: number): string => { @@ -148,6 +154,7 @@ const normalizeUserItem = (item: Partial): User => { user: String(item.user ?? ""), status: Number(item.status ?? 0), flow: Number(item.flow ?? 0), + flowMiB: Number(item.flowMiB ?? 0), num: Number(item.num ?? 0), expTime: item.expTime, flowResetTime: item.flowResetTime ?? 0, @@ -172,6 +179,7 @@ const normalizeUserTunnelItem = (item: Partial): UserTunnel => { tunnelName: String(item.tunnelName ?? ""), status: Number(item.status ?? 0), flow: Number(item.flow ?? 0), + flowMiB: Number(item.flowMiB ?? 0), num: Number(item.num ?? 0), expTime: Number(item.expTime ?? 0), flowResetTime: Number(item.flowResetTime ?? 0), @@ -223,6 +231,8 @@ export default function UserPage() { maxConn: 0, }); const [userFormLoading, setUserFormLoading] = useState(false); + const [userFlowInput, setUserFlowInput] = useState("1000"); + const [userFlowUnit, setUserFlowUnit] = useState("GB"); const [quotaResetLoading, setQuotaResetLoading] = useState(false); const editingUser = useMemo( @@ -263,6 +273,8 @@ export default function UserPage() { onClose: onEditTunnelModalClose, } = useDisclosure(); const [editTunnelForm, setEditTunnelForm] = useState(null); + const [tunnelFlowInput, setTunnelFlowInput] = useState(""); + const [tunnelFlowUnit, setTunnelFlowUnit] = useState("GB"); const [editTunnelLoading, setEditTunnelLoading] = useState(false); // 删除确认相关状态 @@ -506,6 +518,8 @@ export default function UserPage() { const handleAdd = () => { setIsEdit(false); + setUserFlowInput("1000"); + setUserFlowUnit("GB"); setUserForm({ user: "", pwd: "", @@ -524,6 +538,10 @@ export default function UserPage() { const handleEdit = async (user: User) => { setIsEdit(true); + const trafficInput = trafficInputFor(user.flow, user.flowMiB); + + setUserFlowInput(trafficInput.value); + setUserFlowUnit(trafficInput.unit); let currentGroupIds: number[] = []; try { @@ -591,10 +609,21 @@ export default function UserPage() { return; } + const flowMiB = parseTrafficInput(userFlowInput, userFlowUnit); + + if (flowMiB === null) { + toast.error("请输入有效的流量限制,最小单位为 1 MB"); + + return; + } + setUserFormLoading(true); try { const submitData: any = { ...userForm, + flow: Math.ceil(flowMiB / 1024), + flowMiB: + userFlowInput === "99999" && userFlowUnit === "GB" ? 0 : flowMiB, expTime: userForm.expTime.getTime(), groupIds: userForm.groupIds ?? [], }; @@ -756,6 +785,10 @@ export default function UserPage() { }; const handleEditTunnel = (userTunnel: UserTunnel) => { + const trafficInput = trafficInputFor(userTunnel.flow, userTunnel.flowMiB); + + setTunnelFlowInput(trafficInput.value); + setTunnelFlowUnit(trafficInput.unit); setEditTunnelForm({ ...userTunnel, speedId: normalizeSpeedId(userTunnel.speedId), @@ -767,12 +800,24 @@ export default function UserPage() { const handleUpdateTunnel = async () => { if (!editTunnelForm) return; + const flowMiB = parseTrafficInput(tunnelFlowInput, tunnelFlowUnit); + + if (flowMiB === null) { + toast.error("请输入有效的流量限制,最小单位为 1 MB"); + + return; + } + const flow = Math.ceil(flowMiB / 1024); + const storedFlowMiB = + tunnelFlowInput === "99999" && tunnelFlowUnit === "GB" ? 0 : flowMiB; + setEditTunnelLoading(true); try { const speedLimitAutoCleared = isMissingSpeedLimit(editTunnelForm.speedId); const response = await updateUserTunnel({ id: editTunnelForm.id, - flow: editTunnelForm.flow, + flow, + flowMiB: storedFlowMiB, num: editTunnelForm.num, expTime: editTunnelForm.expTime, flowResetTime: editTunnelForm.flowResetTime, @@ -792,6 +837,8 @@ export default function UserPage() { if (currentUser) { const nextTunnel = normalizeUserTunnelItem({ ...editTunnelForm, + flow, + flowMiB: storedFlowMiB, speedId: normalizeSpeedId(editTunnelForm.speedId), speedLimitName: normalizeSpeedId(editTunnelForm.speedId) !== null @@ -1148,7 +1195,7 @@ export default function UserPage() {
限制: - {formatFlow(user.flow, "gb")} + {formatFlowLimit(user.flow, user.flowMiB)}
@@ -1258,9 +1305,9 @@ export default function UserPage() { : null; const usedFlow = calculateUserTotalUsedFlow(user); const flowPercent = - user.flow > 0 + user.flow > 0 && !(user.flow === 99999 && !user.flowMiB) ? Math.min( - (usedFlow / (user.flow * 1024 * 1024 * 1024)) * 100, + (usedFlow / flowLimitBytes(user.flow, user.flowMiB)) * 100, 100, ) : 0; @@ -1316,7 +1363,7 @@ export default function UserPage() {
流量限制 - {formatFlow(user.flow, "gb")} + {formatFlowLimit(user.flow, user.flowMiB)}
@@ -1504,20 +1551,14 @@ export default function UserPage() { setUserForm((prev) => ({ ...prev, pwd: e.target.value })) } /> - { - const value = Math.min( - Math.max(Number(e.target.value) || 0, 1), - 99999, - ); - - setUserForm((prev) => ({ ...prev, flow: value })); + label="流量限制" + unit={userFlowUnit} + value={userFlowInput} + onChange={(value, unit) => { + setUserFlowInput(value); + setUserFlowUnit(unit); }} /> 限制: - {formatFlow(userTunnel.flow, "gb")} + {formatFlowLimit( + userTunnel.flow, + userTunnel.flowMiB, + )}
@@ -2083,21 +2127,14 @@ export default function UserPage() { {editTunnelForm && ( <>
- { - const value = Math.min( - Math.max(Number(e.target.value) || 0, 1), - 99999, - ); - - setEditTunnelForm((prev) => - prev ? { ...prev, flow: value } : null, - ); + { + setTunnelFlowInput(value); + setTunnelFlowUnit(unit); }} /> diff --git a/vite-frontend/src/types/index.ts b/vite-frontend/src/types/index.ts index 3e07e15..d359274 100644 --- a/vite-frontend/src/types/index.ts +++ b/vite-frontend/src/types/index.ts @@ -12,6 +12,7 @@ export interface User { pwd?: string; status: number; // 1-正常, 0-禁用 flow: number; // 流量限制(GB) + flowMiB?: number; // 精确流量限制(MiB),0 表示沿用旧版 GB 字段 num: number; // 转发数量 expTime?: number; // 过期时间戳 flowResetTime?: number; // 流量重置日期(1-31号) @@ -56,6 +57,7 @@ export interface UserTunnel { tunnelName: string; status: number; // 1-正常, 0-禁用 flow: number; // 流量限制(GB) + flowMiB?: number; num: number; // 转发数量 expTime: number; // 过期时间戳 flowResetTime: number; diff --git a/vite-frontend/src/utils/traffic.ts b/vite-frontend/src/utils/traffic.ts new file mode 100644 index 0000000..2f13c1f --- /dev/null +++ b/vite-frontend/src/utils/traffic.ts @@ -0,0 +1,72 @@ +export type TrafficUnit = "MB" | "GB" | "TB" | "PB"; + +export const MIB = 1024 * 1024; +export const GIB = 1024 * MIB; + +export const TRAFFIC_UNIT_MIB: Record = { + MB: 1, + GB: 1024, + TB: 1024 ** 2, + PB: 1024 ** 3, +}; + +const BYTE_UNITS = ["B", "KB", "MB", "GB", "TB", "PB"]; + +export function formatTraffic(bytes: number): string { + if (!Number.isFinite(bytes) || bytes <= 0) return "0 B"; + + let value = bytes; + let unit = 0; + + while (value >= 1024 && unit < BYTE_UNITS.length - 1) { + value /= 1024; + unit++; + } + + return `${unit === 0 ? Math.floor(value) : value.toFixed(2)} ${BYTE_UNITS[unit]}`; +} + +export function flowLimitMiB(flowGB: number, flowMiB?: number): number { + return flowMiB && flowMiB > 0 ? flowMiB : flowGB * 1024; +} + +export function flowLimitBytes(flowGB: number, flowMiB?: number): number { + return flowLimitMiB(flowGB, flowMiB) * MIB; +} + +export function formatFlowLimit(flowGB: number, flowMiB?: number): string { + if (flowGB === 99999 && !flowMiB) return "无限制"; + + return formatTraffic(flowLimitBytes(flowGB, flowMiB)); +} + +export function preferredTrafficUnit(mib: number): TrafficUnit { + if (mib <= 0) return "GB"; + if (mib > 0 && mib % TRAFFIC_UNIT_MIB.PB === 0) return "PB"; + if (mib > 0 && mib % TRAFFIC_UNIT_MIB.TB === 0) return "TB"; + if (mib > 0 && mib % TRAFFIC_UNIT_MIB.GB === 0) return "GB"; + + return "MB"; +} + +export function parseTrafficInput( + value: string, + unit: TrafficUnit, +): number | null { + const amount = Number(value); + const mib = amount * TRAFFIC_UNIT_MIB[unit]; + + if ( + !value.trim() || + !Number.isFinite(amount) || + amount <= 0 || + !Number.isSafeInteger(mib) + ) { + return null; + } + + // Backend stores bytes as int64. Keep the converted value within that range. + if (mib > 8_796_093_022_207) return null; + + return mib; +}