refactor(quota): migrate traffic quota from tunnel to user level

- Replace tunnel_quota table with user_quota table
- Add user-level daily/monthly quota tracking and enforcement
- Update user CRUD to include quota configuration
- Migrate backup/restore to use user quota fields
- Update frontend API and UI for user quota management
This commit is contained in:
sagitchu
2026-03-12 14:17:57 +08:00
parent 30d9552207
commit ad9b336fb9
19 changed files with 648 additions and 615 deletions
@@ -44,10 +44,8 @@ func (h *Handler) processFlowItem(nodeID int64, item flowItem) {
if ok {
inFlow, outFlow := h.scaleFlowByTunnel(forwardID, item.D, item.U)
_ = h.repo.AddFlow(forwardID, userID, userTunnelID, inFlow, outFlow)
if forward, err := h.getForwardRecord(forwardID); err == nil && forward != nil {
if quota, quotaErr := h.repo.AddTunnelQuotaUsage(forward.TunnelID, inFlow+outFlow, time.Now()); quotaErr == nil {
h.enforceTunnelQuotaIfNeeded(forward.TunnelID, quota)
}
if quota, quotaErr := h.repo.AddUserQuotaUsage(userID, inFlow+outFlow, time.Now()); quotaErr == nil {
h.enforceUserQuotaIfNeeded(userID, quota)
}
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
@@ -361,7 +359,7 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
if flowLimit < current {
return errors.New("流量已超额,禁止开启转发")
}
if err := h.ensureTunnelForwardAllowedByQuota(tunnelID, now); err != nil {
if err := h.ensureUserForwardAllowedByQuota(userID, now); err != nil {
return err
}
+1 -1
View File
@@ -101,6 +101,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/user/update", h.userUpdate)
mux.HandleFunc("/api/v1/user/delete", h.userDelete)
mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
mux.HandleFunc("/api/v1/user/quota/reset", h.userQuotaReset)
mux.HandleFunc("/api/v1/user/groups", h.userGroups)
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
@@ -132,7 +133,6 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
mux.HandleFunc("/api/v1/tunnel/update", h.tunnelUpdate)
mux.HandleFunc("/api/v1/tunnel/quota/reset", h.tunnelQuotaReset)
mux.HandleFunc("/api/v1/tunnel/delete", h.tunnelDelete)
mux.HandleFunc("/api/v1/tunnel/diagnose", h.tunnelDiagnose)
mux.HandleFunc("/api/v1/tunnel/diagnose/stream", h.tunnelDiagnoseStream)
+1 -1
View File
@@ -136,7 +136,7 @@ func (h *Handler) runResetAndExpiryJob(now time.Time) {
}
h.resetMonthlyFlow(now)
h.resetTunnelQuotaWindows(now)
h.resetUserQuotaWindows(now)
h.disableExpiredUsers(now.UnixMilli())
h.disableExpiredUserTunnels(now.UnixMilli())
}
+9 -13
View File
@@ -144,7 +144,7 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
}
}
func TestRunResetAndExpiryJobResetsTunnelQuotaAndReEnablesTunnel(t *testing.T) {
func TestRunResetAndExpiryJobResetsUserQuotaAndUnblocksUser(t *testing.T) {
dbPath := filepath.Join(t.TempDir(), "jobs-quota-reset.db")
r, err := repo.Open(dbPath)
if err != nil {
@@ -157,29 +157,25 @@ func TestRunResetAndExpiryJobResetsTunnelQuotaAndReEnablesTunnel(t *testing.T) {
nowMs := now.UnixMilli()
if err := r.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(1, 'quota-reset-tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'quota-reset-user', 'x', 1, 0, 99999, 0, 0, 1, 99999, ?, ?, 1)
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
t.Fatalf("insert user: %v", err)
}
if err := r.DB().Exec(`
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(2, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
`, 11*int64(1024*1024*1024), 11*int64(1024*1024*1024), nowMs, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert tunnel quota: %v", err)
t.Fatalf("insert user quota: %v", err)
}
h.runResetAndExpiryJob(now)
tunnelStatus := mustQueryInt(t, r, `SELECT status FROM tunnel WHERE id = 1`)
if tunnelStatus != 1 {
t.Fatalf("expected tunnel re-enabled after quota reset, got %d", tunnelStatus)
}
dailyUsed := mustQueryInt(t, r, `SELECT daily_used_bytes FROM tunnel_quota WHERE tunnel_id = 1`)
dailyUsed := mustQueryInt(t, r, `SELECT daily_used_bytes FROM user_quota WHERE user_id = 2`)
if dailyUsed != 0 {
t.Fatalf("expected daily quota usage reset, got %d", dailyUsed)
}
quotaDisabled := mustQueryInt(t, r, `SELECT disabled_by_quota FROM tunnel_quota WHERE tunnel_id = 1`)
quotaDisabled := mustQueryInt(t, r, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`)
if quotaDisabled != 0 {
t.Fatalf("expected quota disabled flag cleared, got %d", quotaDisabled)
}
+52 -20
View File
@@ -60,6 +60,12 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
num := asInt(req["num"], 10)
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
flowResetTime := asInt64(req["flowResetTime"], 1)
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
if dailyQuotaGB < 0 || monthlyQuotaGB < 0 {
response.WriteJSON(w, response.ErrDefault("配额不能小于0"))
return
}
roleID := 1
now := time.Now().UnixMilli()
@@ -68,6 +74,22 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if dailyQuotaGB > 0 || monthlyQuotaGB > 0 {
tx := h.repo.BeginTx()
if tx == nil || tx.Error != nil {
response.WriteJSON(w, response.Err(-2, "database unavailable"))
return
}
defer func() { tx.Rollback() }()
if err := h.repo.SaveUserQuotaConfigTx(tx, userID, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit().Error; err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
groupIDs := asInt64Slice(req["groupIds"])
if len(groupIDs) > 0 {
@@ -131,6 +153,8 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
expTime := asInt64(req["expTime"], time.Now().Add(365*24*time.Hour).UnixMilli())
flowResetTime := asInt64(req["flowResetTime"], 1)
status := asInt(req["status"], 1)
_, hasDailyQuota := req["dailyQuotaGB"]
_, hasMonthlyQuota := req["monthlyQuotaGB"]
now := time.Now().UnixMilli()
pwd := asString(req["pwd"])
@@ -147,6 +171,34 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
}
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime)
if hasDailyQuota || hasMonthlyQuota {
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
if !(hasDailyQuota && hasMonthlyQuota) {
if currentQuota, err := h.repo.GetUserQuotaView(id, time.Now()); err == nil && currentQuota != nil {
if !hasDailyQuota {
dailyQuotaGB = currentQuota.DailyLimitGB
}
if !hasMonthlyQuota {
monthlyQuotaGB = currentQuota.MonthlyLimitGB
}
}
}
tx := h.repo.BeginTx()
if tx == nil || tx.Error != nil {
response.WriteJSON(w, response.Err(-2, "database unavailable"))
return
}
defer func() { tx.Rollback() }()
if err := h.repo.SaveUserQuotaConfigTx(tx, id, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit().Error; err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
}
if groupIDsRaw, ok := req["groupIds"]; ok {
newGroupIDs := asInt64Slice(groupIDsRaw)
@@ -485,8 +537,6 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
typeVal := asInt(req["type"], 1)
flow := asInt64(req["flow"], 1)
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
status := asInt(req["status"], 1)
trafficRatio := asFloat(req["trafficRatio"], 1.0)
inIP := asString(req["inIp"])
@@ -582,10 +632,6 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
return
}
tunnelID := tunnel.ID
if err := h.repo.SaveTunnelQuotaConfigTx(tx, tunnelID, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
runtimeState.TunnelID = tunnelID
var federationBindings []repo.FederationTunnelBinding
var federationReleaseRefs []federationRuntimeReleaseRef
@@ -697,8 +743,6 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
now := time.Now().UnixMilli()
typeVal := asInt(req["type"], 1)
dailyQuotaGB := asInt64(req["dailyQuotaGB"], 0)
monthlyQuotaGB := asInt64(req["monthlyQuotaGB"], 0)
ipPreference := asString(req["ipPreference"])
localDomain := h.federationLocalDomain()
@@ -743,10 +787,6 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.repo.SaveTunnelQuotaConfigTx(tx, id, dailyQuotaGB, monthlyQuotaGB, now); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := h.repo.DeleteChainTunnelsByTunnelTx(tx, id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -1360,10 +1400,6 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
return
}
if tunnel.Status != 1 {
if reason, quotaErr := h.tunnelQuotaBlockReason(tunnelID, time.Now().UnixMilli()); quotaErr == nil && reason != "" {
response.WriteJSON(w, response.ErrDefault(reason))
return
}
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法创建转发"))
return
}
@@ -1485,10 +1521,6 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
return
}
if tunnel.Status != 1 {
if reason, quotaErr := h.tunnelQuotaBlockReason(tunnelID, time.Now().UnixMilli()); quotaErr == nil && reason != "" {
response.WriteJSON(w, response.ErrDefault(reason))
return
}
response.WriteJSON(w, response.ErrDefault("隧道已禁用,无法更新转发"))
return
}
@@ -1,140 +0,0 @@
package handler
import (
"errors"
"net/http"
"time"
"go-backend/internal/http/response"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
)
func isTunnelQuotaExceeded(view *model.TunnelQuotaView) bool {
if view == nil {
return false
}
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*bytesPerGB {
return true
}
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*bytesPerGB {
return true
}
return false
}
func (h *Handler) tunnelQuotaBlockReason(tunnelID int64, now int64) (string, error) {
if h == nil || h.repo == nil || tunnelID <= 0 {
return "", nil
}
quota, err := h.repo.GetTunnelQuotaView(tunnelID, time.UnixMilli(now))
if err != nil || quota == nil {
return "", err
}
if quota.DisabledByQuota == 1 || isTunnelQuotaExceeded(quota) {
return "该隧道流量配额已超额,禁止开启转发", nil
}
return "", nil
}
func (h *Handler) enforceTunnelQuotaIfNeeded(tunnelID int64, quota *model.TunnelQuotaView) {
if h == nil || h.repo == nil || tunnelID <= 0 || quota == nil {
return
}
if quota.DisabledByQuota == 1 || !isTunnelQuotaExceeded(quota) {
return
}
forwards, err := h.listForwardsByTunnel(tunnelID)
if err != nil {
return
}
pausedIDs := make([]int64, 0, len(forwards))
now := time.Now().UnixMilli()
for i := range forwards {
if forwards[i].Status != 1 {
continue
}
if err := h.controlForwardServices(&forwards[i], "PauseService", false); err != nil {
continue
}
if err := h.repo.UpdateForwardStatus(forwards[i].ID, 0, now); err != nil {
continue
}
pausedIDs = append(pausedIDs, forwards[i].ID)
}
_ = h.repo.UpdateTunnelStatus(tunnelID, 0, now)
_ = h.repo.MarkTunnelQuotaDisabled(tunnelID, pausedIDs, now)
}
func (h *Handler) applyTunnelQuotaRelease(release *repo.TunnelQuotaRelease, now int64) {
if h == nil || h.repo == nil || release == nil || release.TunnelID <= 0 || !release.EnableTunnel {
return
}
_ = h.repo.UpdateTunnelStatus(release.TunnelID, 1, now)
for _, forwardID := range release.ForwardIDs {
forward, err := h.getForwardRecord(forwardID)
if err != nil || forward == nil {
continue
}
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
continue
}
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
continue
}
_ = h.repo.UpdateForwardStatus(forwardID, 1, now)
}
}
func (h *Handler) resetTunnelQuotaWindows(now time.Time) {
if h == nil || h.repo == nil {
return
}
releases, err := h.repo.RollTunnelQuotaWindows(now)
if err != nil {
return
}
nowMs := now.UnixMilli()
for i := range releases {
h.applyTunnelQuotaRelease(&releases[i], nowMs)
}
}
func (h *Handler) tunnelQuotaReset(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req struct {
TunnelID int64 `json:"tunnelId"`
Scope string `json:"scope"`
}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
if req.TunnelID <= 0 {
response.WriteJSON(w, response.ErrDefault("隧道ID不能为空"))
return
}
release, err := h.repo.ResetTunnelQuotaUsage(req.TunnelID, req.Scope, time.Now())
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
nowMs := time.Now().UnixMilli()
h.applyTunnelQuotaRelease(release, nowMs)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) ensureTunnelForwardAllowedByQuota(tunnelID int64, now int64) error {
reason, err := h.tunnelQuotaBlockReason(tunnelID, now)
if err != nil {
return err
}
if reason != "" {
return errors.New(reason)
}
return nil
}
@@ -0,0 +1,139 @@
package handler
import (
"errors"
"net/http"
"time"
"go-backend/internal/http/response"
"go-backend/internal/store/model"
"go-backend/internal/store/repo"
)
func isUserQuotaExceeded(view *model.UserQuotaView) bool {
if view == nil {
return false
}
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*bytesPerGB {
return true
}
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*bytesPerGB {
return true
}
return false
}
func (h *Handler) userQuotaBlockReason(userID int64, now int64) (string, error) {
if h == nil || h.repo == nil || userID <= 0 {
return "", nil
}
quota, err := h.repo.GetUserQuotaView(userID, time.UnixMilli(now))
if err != nil || quota == nil {
return "", err
}
if quota.DisabledByQuota == 1 || isUserQuotaExceeded(quota) {
return "该用户流量配额已超额,禁止开启转发", nil
}
return "", nil
}
func (h *Handler) enforceUserQuotaIfNeeded(userID int64, quota *model.UserQuotaView) {
if h == nil || h.repo == nil || userID <= 0 || quota == nil {
return
}
if quota.DisabledByQuota == 1 || !isUserQuotaExceeded(quota) {
return
}
forwards, err := h.listActiveForwardsByUser(userID)
if err != nil {
return
}
pausedIDs := make([]int64, 0, len(forwards))
now := time.Now().UnixMilli()
for i := range forwards {
forward := &forwards[i]
if forward.Status != 1 {
continue
}
if err := h.controlForwardServices(forward, "PauseService", false); err != nil {
continue
}
if err := h.repo.UpdateForwardStatus(forward.ID, 0, now); err != nil {
continue
}
pausedIDs = append(pausedIDs, forward.ID)
}
_ = h.repo.MarkUserQuotaDisabled(userID, pausedIDs, now)
}
func (h *Handler) applyUserQuotaRelease(release *repo.UserQuotaRelease, now int64) {
if h == nil || h.repo == nil || release == nil || release.UserID <= 0 || !release.UnblockUser {
return
}
for _, forwardID := range release.ForwardIDs {
forward, err := h.getForwardRecord(forwardID)
if err != nil || forward == nil {
continue
}
if err := h.ensureUserTunnelForwardAllowed(forward.UserID, forward.TunnelID, now); err != nil {
continue
}
if err := h.controlForwardServices(forward, "ResumeService", false); err != nil {
continue
}
_ = h.repo.UpdateForwardStatus(forwardID, 1, now)
}
}
func (h *Handler) resetUserQuotaWindows(now time.Time) {
if h == nil || h.repo == nil {
return
}
releases, err := h.repo.RollUserQuotaWindows(now)
if err != nil {
return
}
nowMs := now.UnixMilli()
for i := range releases {
h.applyUserQuotaRelease(&releases[i], nowMs)
}
}
func (h *Handler) userQuotaReset(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req struct {
UserID int64 `json:"userId"`
Scope string `json:"scope"`
}
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
return
}
if req.UserID <= 0 {
response.WriteJSON(w, response.ErrDefault("用户ID不能为空"))
return
}
release, err := h.repo.ResetUserQuotaUsage(req.UserID, req.Scope, time.Now())
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
nowMs := time.Now().UnixMilli()
h.applyUserQuotaRelease(release, nowMs)
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) ensureUserForwardAllowedByQuota(userID int64, now int64) error {
reason, err := h.userQuotaBlockReason(userID, now)
if err != nil {
return err
}
if reason != "" {
return errors.New(reason)
}
return nil
}
+35 -35
View File
@@ -129,8 +129,8 @@ type Tunnel struct {
func (Tunnel) TableName() string { return "tunnel" }
type TunnelQuota struct {
TunnelID int64 `gorm:"column:tunnel_id;primaryKey"`
type UserQuota struct {
UserID int64 `gorm:"column:user_id;primaryKey"`
DailyLimitGB int64 `gorm:"column:daily_limit_gb;not null;default:0"`
MonthlyLimitGB int64 `gorm:"column:monthly_limit_gb;not null;default:0"`
DailyUsedBytes int64 `gorm:"column:daily_used_bytes;not null;default:0"`
@@ -144,7 +144,7 @@ type TunnelQuota struct {
UpdatedTime int64 `gorm:"column:updated_time;not null"`
}
func (TunnelQuota) TableName() string { return "tunnel_quota" }
func (UserQuota) TableName() string { return "user_quota" }
type ChainTunnel struct {
ID int64 `gorm:"primaryKey;autoIncrement"`
@@ -338,19 +338,23 @@ type BackupData struct {
}
type UserBackup struct {
ID int64 `json:"id"`
User string `json:"user"`
Pwd string `json:"pwd"`
RoleID int `json:"roleId"`
ExpTime int64 `json:"expTime"`
Flow int64 `json:"flow"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
FlowResetTime int64 `json:"flowResetTime"`
Num int `json:"num"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime,omitempty"`
Status int `json:"status"`
ID int64 `json:"id"`
User string `json:"user"`
Pwd string `json:"pwd"`
RoleID int `json:"roleId"`
ExpTime int64 `json:"expTime"`
Flow int64 `json:"flow"`
InFlow int64 `json:"inFlow"`
OutFlow int64 `json:"outFlow"`
FlowResetTime int64 `json:"flowResetTime"`
DailyQuotaGB int64 `json:"dailyQuotaGB,omitempty"`
MonthlyQuotaGB int64 `json:"monthlyQuotaGB,omitempty"`
DisabledByQuota int `json:"disabledByQuota,omitempty"`
QuotaDisabledAt int64 `json:"quotaDisabledAt,omitempty"`
Num int `json:"num"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime,omitempty"`
Status int `json:"status"`
}
type NodeBackup struct {
@@ -383,23 +387,19 @@ type NodeBackup struct {
}
type TunnelBackup struct {
ID int64 `json:"id"`
Name string `json:"name"`
TrafficRatio float64 `json:"trafficRatio"`
Type int `json:"type"`
Protocol string `json:"protocol"`
Flow int64 `json:"flow"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
InIP string `json:"inIp,omitempty"`
Inx int `json:"inx"`
IPPreference string `json:"ipPreference,omitempty"`
DailyQuotaGB int64 `json:"dailyQuotaGB,omitempty"`
MonthlyQuotaGB int64 `json:"monthlyQuotaGB,omitempty"`
DisabledByQuota int `json:"disabledByQuota,omitempty"`
QuotaDisabledAt int64 `json:"quotaDisabledAt,omitempty"`
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
ID int64 `json:"id"`
Name string `json:"name"`
TrafficRatio float64 `json:"trafficRatio"`
Type int `json:"type"`
Protocol string `json:"protocol"`
Flow int64 `json:"flow"`
CreatedTime int64 `json:"createdTime"`
UpdatedTime int64 `json:"updatedTime"`
Status int `json:"status"`
InIP string `json:"inIp,omitempty"`
Inx int `json:"inx"`
IPPreference string `json:"ipPreference,omitempty"`
ChainTunnels []ChainTunnelBackup `json:"chainTunnels,omitempty"`
}
type ChainTunnelBackup struct {
@@ -537,8 +537,8 @@ type TunnelRecord struct {
TrafficRatio float64
}
type TunnelQuotaView struct {
TunnelID int64
type UserQuotaView struct {
UserID int64
DailyLimitGB int64
MonthlyLimitGB int64
DailyUsedBytes int64
+67 -56
View File
@@ -161,13 +161,13 @@ func (r *Repository) Close() error {
func autoMigrateAll(db *gorm.DB) error {
models := []interface{}{
&model.User{},
&model.UserQuota{},
&model.Forward{},
&model.ForwardPort{},
&model.Node{},
&model.SpeedLimit{},
&model.StatisticsFlow{},
&model.Tunnel{},
&model.TunnelQuota{},
&model.ChainTunnel{},
&model.UserTunnel{},
&model.TunnelGroup{},
@@ -664,16 +664,33 @@ func (r *Repository) ListUsers() ([]map[string]interface{}, error) {
if err := r.db.Where("role_id != ?", 0).Order("id DESC").Find(&users).Error; err != nil {
return nil, err
}
userIDs := make([]int64, 0, len(users))
for _, u := range users {
userIDs = append(userIDs, u.ID)
}
quotaMap, err := r.ListUserQuotaViewsByUserIDs(userIDs, time.Now())
if err != nil {
return nil, err
}
items := make([]map[string]interface{}, 0, len(users))
for _, u := range users {
items = append(items, map[string]interface{}{
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,
"flowResetTime": u.FlowResetTime, "createdTime": u.CreatedTime,
"updatedTime": nullableInt64(u.UpdatedTime),
"inFlow": u.InFlow, "outFlow": u.OutFlow,
})
}
if quota := quotaMap[u.ID]; quota != nil {
item["dailyQuotaGB"] = quota.DailyLimitGB
item["monthlyQuotaGB"] = quota.MonthlyLimitGB
item["dailyUsedBytes"] = quota.DailyUsedBytes
item["monthlyUsedBytes"] = quota.MonthlyUsedBytes
item["disabledByQuota"] = quota.DisabledByQuota
item["quotaDisabledAt"] = quota.DisabledAt
}
items = append(items, item)
}
return items, nil
}
@@ -958,7 +975,6 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
tunnelMap := make(map[int64]map[string]interface{})
orderedIDs := make([]int64, 0, len(tunnels))
tunnelIDs := make([]int64, 0, len(tunnels))
for _, t := range tunnels {
tunnelMap[t.ID] = map[string]interface{}{
@@ -972,24 +988,6 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) {
"chainNodes": make([][]map[string]interface{}, 0),
}
orderedIDs = append(orderedIDs, t.ID)
tunnelIDs = append(tunnelIDs, t.ID)
}
quotaMap, err := r.ListTunnelQuotaViewsByTunnelIDs(tunnelIDs, time.Now())
if err != nil {
return nil, err
}
for tunnelID, quota := range quotaMap {
item := tunnelMap[tunnelID]
if item == nil || quota == nil {
continue
}
item["dailyQuotaGB"] = quota.DailyLimitGB
item["monthlyQuotaGB"] = quota.MonthlyLimitGB
item["dailyUsedBytes"] = quota.DailyUsedBytes
item["monthlyUsedBytes"] = quota.MonthlyUsedBytes
item["disabledByQuota"] = quota.DisabledByQuota
item["quotaDisabledAt"] = quota.DisabledAt
}
// Build node IP map
@@ -1814,6 +1812,14 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
if err := r.db.Order("id ASC").Find(&users).Error; err != nil {
return nil, err
}
userIDs := make([]int64, 0, len(users))
for _, u := range users {
userIDs = append(userIDs, u.ID)
}
quotaMap, err := r.ListUserQuotaViewsByUserIDs(userIDs, time.Now())
if err != nil {
return nil, err
}
out := make([]model.UserBackup, 0, len(users))
for _, u := range users {
b := model.UserBackup{
@@ -1822,6 +1828,12 @@ func (r *Repository) exportUsers() ([]model.UserBackup, error) {
FlowResetTime: u.FlowResetTime, Num: u.Num,
CreatedTime: u.CreatedTime, Status: u.Status,
}
if quota := quotaMap[u.ID]; quota != nil {
b.DailyQuotaGB = quota.DailyLimitGB
b.MonthlyQuotaGB = quota.MonthlyLimitGB
b.DisabledByQuota = quota.DisabledByQuota
b.QuotaDisabledAt = quota.DisabledAt
}
if u.UpdatedTime.Valid {
b.UpdatedTime = u.UpdatedTime.Int64
}
@@ -1882,14 +1894,6 @@ func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) {
if err := r.db.Order("inx ASC, id ASC").Find(&tunnels).Error; err != nil {
return nil, err
}
quotaIDs := make([]int64, 0, len(tunnels))
for _, t := range tunnels {
quotaIDs = append(quotaIDs, t.ID)
}
quotaMap, err := r.ListTunnelQuotaViewsByTunnelIDs(quotaIDs, time.Now())
if err != nil {
return nil, err
}
out := make([]model.TunnelBackup, 0, len(tunnels))
for _, t := range tunnels {
b := model.TunnelBackup{
@@ -1898,12 +1902,6 @@ func (r *Repository) exportTunnels() ([]model.TunnelBackup, error) {
CreatedTime: t.CreatedTime, UpdatedTime: t.UpdatedTime,
Status: t.Status, Inx: t.Inx, IPPreference: t.IPPreference,
}
if quota := quotaMap[t.ID]; quota != nil {
b.DailyQuotaGB = quota.DailyLimitGB
b.MonthlyQuotaGB = quota.MonthlyLimitGB
b.DisabledByQuota = quota.DisabledByQuota
b.QuotaDisabledAt = quota.DisabledAt
}
if t.InIP.Valid {
b.InIP = t.InIP.String
}
@@ -2200,6 +2198,39 @@ func importUsers(tx *gorm.DB, users []model.UserBackup, now int64) (int, error)
if err != nil {
return count, err
}
if u.DailyQuotaGB > 0 || u.MonthlyQuotaGB > 0 || u.DisabledByQuota != 0 || u.QuotaDisabledAt > 0 {
current := time.UnixMilli(now)
dayKey := int64(current.Year()*10000 + int(current.Month())*100 + current.Day())
monthKey := int64(current.Year()*100 + int(current.Month()))
quotaItem := model.UserQuota{
UserID: u.ID,
DailyLimitGB: u.DailyQuotaGB,
MonthlyLimitGB: u.MonthlyQuotaGB,
DailyUsedBytes: 0,
MonthlyUsedBytes: 0,
DayKey: dayKey,
MonthKey: monthKey,
DisabledByQuota: u.DisabledByQuota,
DisabledAt: u.QuotaDisabledAt,
PausedForwardIDs: "",
CreatedTime: now,
UpdatedTime: now,
}
err = tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "user_id"}},
DoUpdates: clause.AssignmentColumns([]string{
"daily_limit_gb", "monthly_limit_gb", "daily_used_bytes", "monthly_used_bytes",
"day_key", "month_key", "disabled_by_quota", "disabled_at", "paused_forward_ids", "updated_time",
}),
}).Create(&quotaItem).Error
if err != nil {
return count, err
}
} else {
if err := tx.Where("user_id = ?", u.ID).Delete(&model.UserQuota{}).Error; err != nil {
return count, err
}
}
count++
}
return count, nil
@@ -2277,26 +2308,6 @@ func importTunnels(tx *gorm.DB, tunnels []model.TunnelBackup, now int64) (int, e
if err != nil {
return count, err
}
quotaItem := model.TunnelQuota{
TunnelID: t.ID,
DailyLimitGB: t.DailyQuotaGB,
MonthlyLimitGB: t.MonthlyQuotaGB,
DisabledByQuota: t.DisabledByQuota,
DisabledAt: t.QuotaDisabledAt,
DayKey: int64(time.Now().Year()*10000 + int(time.Now().Month())*100 + time.Now().Day()),
MonthKey: int64(time.Now().Year()*100 + int(time.Now().Month())),
CreatedTime: now,
UpdatedTime: now,
}
err = tx.Clauses(clause.OnConflict{
Columns: []clause.Column{{Name: "tunnel_id"}},
DoUpdates: clause.AssignmentColumns([]string{
"daily_limit_gb", "monthly_limit_gb", "disabled_by_quota", "disabled_at", "updated_time",
}),
}).Create(&quotaItem).Error
if err != nil {
return count, err
}
for _, ct := range t.ChainTunnels {
chainItem := model.ChainTunnel{
ID: ct.ID,
@@ -145,6 +145,9 @@ func (r *Repository) DeleteUserCascade(userID int64) error {
if err := tx.Where("user_id = ?", userID).Delete(&model.StatisticsFlow{}).Error; err != nil {
return err
}
if err := tx.Where("user_id = ?", userID).Delete(&model.UserQuota{}).Error; err != nil {
return err
}
return tx.Where("id = ?", userID).Delete(&model.User{}).Error
})
}
@@ -13,21 +13,21 @@ import (
"gorm.io/gorm/clause"
)
const tunnelQuotaBytesPerGB int64 = 1024 * 1024 * 1024
const userQuotaBytesPerGB int64 = 1024 * 1024 * 1024
type TunnelQuotaRelease struct {
TunnelID int64
ForwardIDs []int64
EnableTunnel bool
type UserQuotaRelease struct {
UserID int64
ForwardIDs []int64
UnblockUser bool
}
func tunnelQuotaWindowKeys(now time.Time) (int64, int64) {
func userQuotaWindowKeys(now time.Time) (int64, int64) {
return int64(now.Year()*10000 + int(now.Month())*100 + now.Day()), int64(now.Year()*100 + int(now.Month()))
}
func cloneTunnelQuotaView(q model.TunnelQuota) *model.TunnelQuotaView {
return &model.TunnelQuotaView{
TunnelID: q.TunnelID,
func cloneUserQuotaView(q model.UserQuota) *model.UserQuotaView {
return &model.UserQuotaView{
UserID: q.UserID,
DailyLimitGB: q.DailyLimitGB,
MonthlyLimitGB: q.MonthlyLimitGB,
DailyUsedBytes: q.DailyUsedBytes,
@@ -40,11 +40,11 @@ func cloneTunnelQuotaView(q model.TunnelQuota) *model.TunnelQuotaView {
}
}
func normalizeTunnelQuotaView(view *model.TunnelQuotaView, now time.Time) *model.TunnelQuotaView {
func normalizeUserQuotaView(view *model.UserQuotaView, now time.Time) *model.UserQuotaView {
if view == nil {
return nil
}
dayKey, monthKey := tunnelQuotaWindowKeys(now)
dayKey, monthKey := userQuotaWindowKeys(now)
out := *view
if out.DayKey != dayKey {
out.DayKey = dayKey
@@ -57,14 +57,14 @@ func normalizeTunnelQuotaView(view *model.TunnelQuotaView, now time.Time) *model
return &out
}
func tunnelQuotaExceeded(view *model.TunnelQuotaView) bool {
func userQuotaExceeded(view *model.UserQuotaView) bool {
if view == nil {
return false
}
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*tunnelQuotaBytesPerGB {
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*userQuotaBytesPerGB {
return true
}
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*tunnelQuotaBytesPerGB {
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*userQuotaBytesPerGB {
return true
}
return false
@@ -107,13 +107,13 @@ func joinPausedForwardIDs(ids []int64) string {
return strings.Join(parts, ",")
}
func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now time.Time) (*model.TunnelQuota, error) {
func (r *Repository) loadOrCreateUserQuotaTx(tx *gorm.DB, userID int64, now time.Time) (*model.UserQuota, error) {
if tx == nil {
return nil, errors.New("database unavailable")
}
dayKey, monthKey := tunnelQuotaWindowKeys(now)
q := &model.TunnelQuota{}
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("tunnel_id = ?", tunnelID).First(q).Error
dayKey, monthKey := userQuotaWindowKeys(now)
q := &model.UserQuota{}
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Where("user_id = ?", userID).First(q).Error
if err == nil {
return q, nil
}
@@ -121,8 +121,8 @@ func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now
return nil, err
}
nowMs := now.UnixMilli()
q = &model.TunnelQuota{
TunnelID: tunnelID,
q = &model.UserQuota{
UserID: userID,
DayKey: dayKey,
MonthKey: monthKey,
CreatedTime: nowMs,
@@ -135,12 +135,12 @@ func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now
return q, nil
}
func applyTunnelQuotaWindowRoll(q *model.TunnelQuota, now time.Time) bool {
func applyUserQuotaWindowRoll(q *model.UserQuota, now time.Time) bool {
if q == nil {
return false
}
changed := false
dayKey, monthKey := tunnelQuotaWindowKeys(now)
dayKey, monthKey := userQuotaWindowKeys(now)
if q.DayKey != dayKey {
q.DayKey = dayKey
q.DailyUsedBytes = 0
@@ -154,18 +154,18 @@ func applyTunnelQuotaWindowRoll(q *model.TunnelQuota, now time.Time) bool {
return changed
}
func (r *Repository) SaveTunnelQuotaConfigTx(tx *gorm.DB, tunnelID, dailyLimitGB, monthlyLimitGB int64, now int64) error {
func (r *Repository) SaveUserQuotaConfigTx(tx *gorm.DB, userID, dailyLimitGB, monthlyLimitGB int64, now int64) error {
if tx == nil {
return errors.New("database unavailable")
}
if tunnelID <= 0 {
return errors.New("tunnel id is required")
if userID <= 0 {
return errors.New("user id is required")
}
if dailyLimitGB < 0 || monthlyLimitGB < 0 {
return errors.New("quota limit cannot be negative")
}
current := time.UnixMilli(now)
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, current)
q, err := r.loadOrCreateUserQuotaTx(tx, userID, current)
if err != nil {
return err
}
@@ -175,69 +175,69 @@ func (r *Repository) SaveTunnelQuotaConfigTx(tx *gorm.DB, tunnelID, dailyLimitGB
"updated_time": now,
}
if q.DayKey == 0 || q.MonthKey == 0 {
dayKey, monthKey := tunnelQuotaWindowKeys(current)
dayKey, monthKey := userQuotaWindowKeys(current)
updates["day_key"] = dayKey
updates["month_key"] = monthKey
}
return tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(updates).Error
return tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(updates).Error
}
func (r *Repository) ListTunnelQuotaViewsByTunnelIDs(tunnelIDs []int64, now time.Time) (map[int64]*model.TunnelQuotaView, error) {
func (r *Repository) ListUserQuotaViewsByUserIDs(userIDs []int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
out := make(map[int64]*model.TunnelQuotaView)
if len(tunnelIDs) == 0 {
out := make(map[int64]*model.UserQuotaView)
if len(userIDs) == 0 {
return out, nil
}
var rows []model.TunnelQuota
if err := r.db.Where("tunnel_id IN ?", tunnelIDs).Find(&rows).Error; err != nil {
var rows []model.UserQuota
if err := r.db.Where("user_id IN ?", userIDs).Find(&rows).Error; err != nil {
return nil, err
}
for _, row := range rows {
out[row.TunnelID] = normalizeTunnelQuotaView(cloneTunnelQuotaView(row), now)
out[row.UserID] = normalizeUserQuotaView(cloneUserQuotaView(row), now)
}
return out, nil
}
func (r *Repository) GetTunnelQuotaView(tunnelID int64, now time.Time) (*model.TunnelQuotaView, error) {
func (r *Repository) GetUserQuotaView(userID int64, now time.Time) (*model.UserQuotaView, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if tunnelID <= 0 {
if userID <= 0 {
return nil, nil
}
var row model.TunnelQuota
err := r.db.Where("tunnel_id = ?", tunnelID).First(&row).Error
var row model.UserQuota
err := r.db.Where("user_id = ?", userID).First(&row).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
return normalizeTunnelQuotaView(cloneTunnelQuotaView(row), now), nil
return normalizeUserQuotaView(cloneUserQuotaView(row), now), nil
}
func (r *Repository) AddTunnelQuotaUsage(tunnelID int64, usedBytes int64, now time.Time) (*model.TunnelQuotaView, error) {
func (r *Repository) AddUserQuotaUsage(userID int64, usedBytes int64, now time.Time) (*model.UserQuotaView, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if tunnelID <= 0 {
if userID <= 0 {
return nil, nil
}
result := &model.TunnelQuotaView{}
result := &model.UserQuotaView{}
err := r.db.Transaction(func(tx *gorm.DB) error {
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, now)
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
if err != nil {
return err
}
applyTunnelQuotaWindowRoll(q, now)
applyUserQuotaWindowRoll(q, now)
if usedBytes > 0 {
q.DailyUsedBytes += usedBytes
q.MonthlyUsedBytes += usedBytes
}
q.UpdatedTime = now.UnixMilli()
if err := tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
"daily_used_bytes": q.DailyUsedBytes,
"monthly_used_bytes": q.MonthlyUsedBytes,
"day_key": q.DayKey,
@@ -246,23 +246,23 @@ func (r *Repository) AddTunnelQuotaUsage(tunnelID int64, usedBytes int64, now ti
}).Error; err != nil {
return err
}
*result = *cloneTunnelQuotaView(*q)
*result = *cloneUserQuotaView(*q)
return nil
})
if err != nil {
return nil, err
}
return normalizeTunnelQuotaView(result, now), nil
return normalizeUserQuotaView(result, now), nil
}
func (r *Repository) MarkTunnelQuotaDisabled(tunnelID int64, pausedForwardIDs []int64, now int64) error {
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if tunnelID <= 0 {
return errors.New("tunnel id is required")
if userID <= 0 {
return errors.New("user id is required")
}
return r.db.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
return r.db.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
"disabled_by_quota": 1,
"disabled_at": now,
"paused_forward_ids": joinPausedForwardIDs(pausedForwardIDs),
@@ -270,12 +270,12 @@ func (r *Repository) MarkTunnelQuotaDisabled(tunnelID int64, pausedForwardIDs []
}).Error
}
func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now time.Time) (*TunnelQuotaRelease, error) {
func (r *Repository) ResetUserQuotaUsage(userID int64, scope string, now time.Time) (*UserQuotaRelease, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
if tunnelID <= 0 {
return nil, errors.New("tunnel id is required")
if userID <= 0 {
return nil, errors.New("user id is required")
}
scope = strings.TrimSpace(strings.ToLower(scope))
if scope == "" {
@@ -284,13 +284,13 @@ func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now tim
if scope != "daily" && scope != "monthly" && scope != "all" {
return nil, fmt.Errorf("unsupported quota reset scope: %s", scope)
}
var release *TunnelQuotaRelease
var release *UserQuotaRelease
err := r.db.Transaction(func(tx *gorm.DB) error {
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, now)
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
if err != nil {
return err
}
applyTunnelQuotaWindowRoll(q, now)
applyUserQuotaWindowRoll(q, now)
switch scope {
case "daily":
q.DailyUsedBytes = 0
@@ -301,15 +301,15 @@ func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now tim
q.MonthlyUsedBytes = 0
}
q.UpdatedTime = now.UnixMilli()
release = &TunnelQuotaRelease{TunnelID: tunnelID}
if q.DisabledByQuota == 1 && !tunnelQuotaExceeded(cloneTunnelQuotaView(*q)) {
release.EnableTunnel = true
release = &UserQuotaRelease{UserID: userID}
if q.DisabledByQuota == 1 && !userQuotaExceeded(cloneUserQuotaView(*q)) {
release.UnblockUser = true
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
q.DisabledByQuota = 0
q.DisabledAt = 0
q.PausedForwardIDs = ""
}
return tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
return tx.Model(&model.UserQuota{}).Where("user_id = ?", userID).Updates(map[string]interface{}{
"daily_used_bytes": q.DailyUsedBytes,
"monthly_used_bytes": q.MonthlyUsedBytes,
"day_key": q.DayKey,
@@ -326,23 +326,23 @@ func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now tim
return release, nil
}
func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease, error) {
func (r *Repository) RollUserQuotaWindows(now time.Time) ([]UserQuotaRelease, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var releases []TunnelQuotaRelease
var releases []UserQuotaRelease
err := r.db.Transaction(func(tx *gorm.DB) error {
var rows []model.TunnelQuota
var rows []model.UserQuota
if err := tx.Find(&rows).Error; err != nil {
return err
}
nowMs := now.UnixMilli()
for _, row := range rows {
q := row
changed := applyTunnelQuotaWindowRoll(&q, now)
release := TunnelQuotaRelease{TunnelID: q.TunnelID}
if q.DisabledByQuota == 1 && !tunnelQuotaExceeded(cloneTunnelQuotaView(q)) {
release.EnableTunnel = true
changed := applyUserQuotaWindowRoll(&q, now)
release := UserQuotaRelease{UserID: q.UserID}
if q.DisabledByQuota == 1 && !userQuotaExceeded(cloneUserQuotaView(q)) {
release.UnblockUser = true
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
q.DisabledByQuota = 0
q.DisabledAt = 0
@@ -353,7 +353,7 @@ func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease
continue
}
q.UpdatedTime = nowMs
if err := tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", q.TunnelID).Updates(map[string]interface{}{
if err := tx.Model(&model.UserQuota{}).Where("user_id = ?", q.UserID).Updates(map[string]interface{}{
"daily_used_bytes": q.DailyUsedBytes,
"monthly_used_bytes": q.MonthlyUsedBytes,
"day_key": q.DayKey,
@@ -365,7 +365,7 @@ func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease
}).Error; err != nil {
return err
}
if release.EnableTunnel {
if release.UnblockUser {
releases = append(releases, release)
}
}
@@ -376,16 +376,3 @@ func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease
}
return releases, nil
}
func (r *Repository) UpdateTunnelStatus(tunnelID int64, status int, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if tunnelID <= 0 {
return errors.New("tunnel id is required")
}
return r.db.Model(&model.Tunnel{}).Where("id = ?", tunnelID).Updates(map[string]interface{}{
"status": status,
"updated_time": now,
}).Error
}
@@ -13,21 +13,24 @@ import (
"go-backend/internal/http/response"
)
func TestForwardCreateBlockedWhenTunnelQuotaExceeded(t *testing.T) {
func TestForwardCreateBlockedWhenUserQuotaExceeded(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now().UnixMilli()
now := time.Now()
nowMs := now.UnixMilli()
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
monthKey := int64(now.Year()*100 + int(now.Month()))
if err := repo.DB().Exec(`
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'quota_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
`, now, now).Error; err != nil {
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(1, 'quota_tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
`, now, now).Error; err != nil {
VALUES(1, 'quota_tunnel', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := repo.DB().Exec(`
@@ -37,10 +40,10 @@ func TestForwardCreateBlockedWhenTunnelQuotaExceeded(t *testing.T) {
t.Fatalf("insert user_tunnel: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
`, 11*contractBytesPerGB, 11*contractBytesPerGB, now, now, now).Error; err != nil {
t.Fatalf("insert tunnel_quota: %v", err)
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '', ?, ?)
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user_quota: %v", err)
}
token, err := auth.GenerateToken(2, "quota_user", 1, secret)
@@ -59,28 +62,31 @@ func TestForwardCreateBlockedWhenTunnelQuotaExceeded(t *testing.T) {
t.Fatalf("decode response: %v", err)
}
if out.Code == 0 {
t.Fatalf("expected non-zero code when tunnel quota exceeded")
t.Fatalf("expected non-zero code when user quota exceeded")
}
if !strings.Contains(out.Msg, "配额") {
t.Fatalf("expected quota error, got %q", out.Msg)
}
}
func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
func TestForwardResumeBlockedWhenUserQuotaExceeded(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now().UnixMilli()
now := time.Now()
nowMs := now.UnixMilli()
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
monthKey := int64(now.Year()*100 + int(now.Month()))
if err := repo.DB().Exec(`
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'quota_resume_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
`, now, now).Error; err != nil {
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(1, 'quota_resume_tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
`, now, now).Error; err != nil {
VALUES(1, 'quota_resume_tunnel', 1.0, 1, 'tls', 1, ?, ?, 1, NULL, 0)
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := repo.DB().Exec(`
@@ -92,14 +98,14 @@ func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
if err := repo.DB().Exec(`
INSERT INTO forward(id, user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
VALUES(1, 2, 'quota_resume_user', 'quota_resume_forward', 1, '1.1.1.1:53', 'fifo', 0, 0, ?, ?, 0, 0)
`, now, now).Error; err != nil {
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '1', ?, ?)
`, 11*contractBytesPerGB, 11*contractBytesPerGB, now, now, now).Error; err != nil {
t.Fatalf("insert tunnel_quota: %v", err)
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '1', ?, ?)
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user_quota: %v", err)
}
token, err := auth.GenerateToken(2, "quota_resume_user", 1, secret)
@@ -118,7 +124,7 @@ func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
t.Fatalf("decode response: %v", err)
}
if out.Code == 0 {
t.Fatalf("expected non-zero code when tunnel quota exceeded")
t.Fatalf("expected non-zero code when user quota exceeded")
}
if !strings.Contains(out.Msg, "配额") {
t.Fatalf("expected quota error, got %q", out.Msg)
@@ -129,29 +135,32 @@ func TestForwardResumeBlockedWhenTunnelQuotaExceeded(t *testing.T) {
}
}
func TestTunnelQuotaResetReEnablesTunnel(t *testing.T) {
func TestUserQuotaResetClearsDisableFlag(t *testing.T) {
secret := "contract-jwt-secret"
router, repo := setupContractRouter(t, secret)
now := time.Now().UnixMilli()
now := time.Now()
nowMs := now.UnixMilli()
dayKey := int64(now.Year()*10000 + int(now.Month())*100 + now.Day())
monthKey := int64(now.Year()*100 + int(now.Month()))
if err := repo.DB().Exec(`
INSERT INTO tunnel(id, name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
VALUES(1, 'quota_reset_tunnel', 1.0, 1, 'tls', 1, ?, ?, 0, NULL, 0)
`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
VALUES(2, 'quota_reset_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
`, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
if err := repo.DB().Exec(`
INSERT INTO tunnel_quota(tunnel_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(1, 10, 0, ?, ?, 20260311, 202603, 1, ?, '', ?, ?)
`, 11*contractBytesPerGB, 11*contractBytesPerGB, now, now, now).Error; err != nil {
t.Fatalf("insert tunnel_quota: %v", err)
INSERT INTO user_quota(user_id, daily_limit_gb, monthly_limit_gb, daily_used_bytes, monthly_used_bytes, day_key, month_key, disabled_by_quota, disabled_at, paused_forward_ids, created_time, updated_time)
VALUES(2, 10, 0, ?, ?, ?, ?, 1, ?, '', ?, ?)
`, 11*contractBytesPerGB, 11*contractBytesPerGB, dayKey, monthKey, nowMs, nowMs, nowMs).Error; err != nil {
t.Fatalf("insert user_quota: %v", err)
}
token, err := auth.GenerateToken(1, "admin", 0, secret)
if err != nil {
t.Fatalf("generate token: %v", err)
}
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/quota/reset", bytes.NewBufferString(`{"tunnelId":1,"scope":"all"}`))
req := httptest.NewRequest(http.MethodPost, "/api/v1/user/quota/reset", bytes.NewBufferString(`{"userId":2,"scope":"all"}`))
req.Header.Set("Authorization", token)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
@@ -165,11 +174,7 @@ func TestTunnelQuotaResetReEnablesTunnel(t *testing.T) {
if out.Code != 0 {
t.Fatalf("expected reset success, got code=%d msg=%q", out.Code, out.Msg)
}
tunnelStatus := mustQueryInt(t, repo, `SELECT status FROM tunnel WHERE id = 1`)
if tunnelStatus != 1 {
t.Fatalf("expected tunnel to be re-enabled, got %d", tunnelStatus)
}
quotaDisabled := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM tunnel_quota WHERE tunnel_id = 1`)
quotaDisabled := mustQueryInt(t, repo, `SELECT disabled_by_quota FROM user_quota WHERE user_id = 2`)
if quotaDisabled != 0 {
t.Fatalf("expected quota disable flag cleared, got %d", quotaDisabled)
}