mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-09 03:06:37 +08:00
Implement per-tunnel traffic quota feature: - Add TunnelQuota model with daily/monthly usage tracking - Integrate quota enforcement into flow accumulation path - Pause forwards and disable tunnel when quota exceeded - Block new forward creation/resume when tunnel quota disabled - Auto-reset daily/monthly windows at 00:05 via maintenance job - Add manual reset API endpoint for admins - Include quota config in tunnel backup/restore - Add frontend UI for quota settings and usage display Entire-Checkpoint: e629b27ca437
This commit is contained in:
@@ -44,6 +44,11 @@ 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)
|
||||
}
|
||||
}
|
||||
h.processPeerShareFlowFromForward(forwardID, nodeID, serviceName, item)
|
||||
|
||||
if userTunnelID > 0 {
|
||||
@@ -356,6 +361,9 @@ func (h *Handler) ensureUserTunnelForwardAllowed(userID int64, tunnelID int64, n
|
||||
if flowLimit < current {
|
||||
return errors.New("流量已超额,禁止开启转发")
|
||||
}
|
||||
if err := h.ensureTunnelForwardAllowedByQuota(tunnelID, now); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(userID, tunnelID)
|
||||
if err != nil {
|
||||
|
||||
@@ -132,6 +132,7 @@ 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)
|
||||
|
||||
@@ -136,6 +136,7 @@ func (h *Handler) runResetAndExpiryJob(now time.Time) {
|
||||
}
|
||||
|
||||
h.resetMonthlyFlow(now)
|
||||
h.resetTunnelQuotaWindows(now)
|
||||
h.disableExpiredUsers(now.UnixMilli())
|
||||
h.disableExpiredUserTunnels(now.UnixMilli())
|
||||
}
|
||||
|
||||
@@ -143,3 +143,44 @@ func TestRunResetAndExpiryJobResetsFlowAndDisablesExpiredRecords(t *testing.T) {
|
||||
t.Fatalf("expected non-expiring forward to remain enabled, got status=%d", nonExpForwardStatus)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunResetAndExpiryJobResetsTunnelQuotaAndReEnablesTunnel(t *testing.T) {
|
||||
dbPath := filepath.Join(t.TempDir(), "jobs-quota-reset.db")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open sqlite: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = r.Close() })
|
||||
|
||||
h := New(r, "secret")
|
||||
now := time.Date(2026, 3, 12, 0, 0, 5, 0, time.UTC)
|
||||
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)
|
||||
`, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %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, ?, '', ?, ?)
|
||||
`, 11*int64(1024*1024*1024), 11*int64(1024*1024*1024), nowMs, nowMs, nowMs).Error; err != nil {
|
||||
t.Fatalf("insert tunnel 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`)
|
||||
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`)
|
||||
if quotaDisabled != 0 {
|
||||
t.Fatalf("expected quota disabled flag cleared, got %d", quotaDisabled)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -485,6 +485,8 @@ 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"])
|
||||
@@ -580,6 +582,10 @@ 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
|
||||
@@ -691,6 +697,8 @@ 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()
|
||||
|
||||
@@ -735,6 +743,10 @@ 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()))
|
||||
@@ -1299,6 +1311,10 @@ 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
|
||||
}
|
||||
@@ -1420,6 +1436,10 @@ 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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
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
|
||||
}
|
||||
@@ -129,6 +129,23 @@ type Tunnel struct {
|
||||
|
||||
func (Tunnel) TableName() string { return "tunnel" }
|
||||
|
||||
type TunnelQuota struct {
|
||||
TunnelID int64 `gorm:"column:tunnel_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"`
|
||||
MonthlyUsedBytes int64 `gorm:"column:monthly_used_bytes;not null;default:0"`
|
||||
DayKey int64 `gorm:"column:day_key;not null;default:0"`
|
||||
MonthKey int64 `gorm:"column:month_key;not null;default:0"`
|
||||
DisabledByQuota int `gorm:"column:disabled_by_quota;not null;default:0"`
|
||||
DisabledAt int64 `gorm:"column:disabled_at;not null;default:0"`
|
||||
PausedForwardIDs string `gorm:"column:paused_forward_ids;type:text;not null;default:''"`
|
||||
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||
}
|
||||
|
||||
func (TunnelQuota) TableName() string { return "tunnel_quota" }
|
||||
|
||||
type ChainTunnel struct {
|
||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||
TunnelID int64 `gorm:"column:tunnel_id;not null"`
|
||||
@@ -366,19 +383,23 @@ 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"`
|
||||
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"`
|
||||
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"`
|
||||
}
|
||||
|
||||
type ChainTunnelBackup struct {
|
||||
@@ -516,6 +537,19 @@ type TunnelRecord struct {
|
||||
TrafficRatio float64
|
||||
}
|
||||
|
||||
type TunnelQuotaView struct {
|
||||
TunnelID int64
|
||||
DailyLimitGB int64
|
||||
MonthlyLimitGB int64
|
||||
DailyUsedBytes int64
|
||||
MonthlyUsedBytes int64
|
||||
DayKey int64
|
||||
MonthKey int64
|
||||
DisabledByQuota int
|
||||
DisabledAt int64
|
||||
PausedForwardIDs string
|
||||
}
|
||||
|
||||
// ForwardPortRecord is a forward port mapping used by control plane.
|
||||
type ForwardPortRecord struct {
|
||||
NodeID int64
|
||||
|
||||
@@ -167,6 +167,7 @@ func autoMigrateAll(db *gorm.DB) error {
|
||||
&model.SpeedLimit{},
|
||||
&model.StatisticsFlow{},
|
||||
&model.Tunnel{},
|
||||
&model.TunnelQuota{},
|
||||
&model.ChainTunnel{},
|
||||
&model.UserTunnel{},
|
||||
&model.TunnelGroup{},
|
||||
@@ -957,6 +958,7 @@ 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{}{
|
||||
@@ -970,6 +972,24 @@ 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
|
||||
@@ -1862,6 +1882,14 @@ 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{
|
||||
@@ -1870,6 +1898,12 @@ 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
|
||||
}
|
||||
@@ -2243,6 +2277,26 @@ 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("aItem).Error
|
||||
if err != nil {
|
||||
return count, err
|
||||
}
|
||||
for _, ct := range t.ChainTunnels {
|
||||
chainItem := model.ChainTunnel{
|
||||
ID: ct.ID,
|
||||
|
||||
@@ -0,0 +1,391 @@
|
||||
package repo
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/store/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
const tunnelQuotaBytesPerGB int64 = 1024 * 1024 * 1024
|
||||
|
||||
type TunnelQuotaRelease struct {
|
||||
TunnelID int64
|
||||
ForwardIDs []int64
|
||||
EnableTunnel bool
|
||||
}
|
||||
|
||||
func tunnelQuotaWindowKeys(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,
|
||||
DailyLimitGB: q.DailyLimitGB,
|
||||
MonthlyLimitGB: q.MonthlyLimitGB,
|
||||
DailyUsedBytes: q.DailyUsedBytes,
|
||||
MonthlyUsedBytes: q.MonthlyUsedBytes,
|
||||
DayKey: q.DayKey,
|
||||
MonthKey: q.MonthKey,
|
||||
DisabledByQuota: q.DisabledByQuota,
|
||||
DisabledAt: q.DisabledAt,
|
||||
PausedForwardIDs: q.PausedForwardIDs,
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeTunnelQuotaView(view *model.TunnelQuotaView, now time.Time) *model.TunnelQuotaView {
|
||||
if view == nil {
|
||||
return nil
|
||||
}
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(now)
|
||||
out := *view
|
||||
if out.DayKey != dayKey {
|
||||
out.DayKey = dayKey
|
||||
out.DailyUsedBytes = 0
|
||||
}
|
||||
if out.MonthKey != monthKey {
|
||||
out.MonthKey = monthKey
|
||||
out.MonthlyUsedBytes = 0
|
||||
}
|
||||
return &out
|
||||
}
|
||||
|
||||
func tunnelQuotaExceeded(view *model.TunnelQuotaView) bool {
|
||||
if view == nil {
|
||||
return false
|
||||
}
|
||||
if view.DailyLimitGB > 0 && view.DailyUsedBytes >= view.DailyLimitGB*tunnelQuotaBytesPerGB {
|
||||
return true
|
||||
}
|
||||
if view.MonthlyLimitGB > 0 && view.MonthlyUsedBytes >= view.MonthlyLimitGB*tunnelQuotaBytesPerGB {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func parsePausedForwardIDs(raw string) []int64 {
|
||||
parts := strings.Split(strings.TrimSpace(raw), ",")
|
||||
out := make([]int64, 0, len(parts))
|
||||
seen := make(map[int64]struct{}, len(parts))
|
||||
for _, part := range parts {
|
||||
id, err := strconv.ParseInt(strings.TrimSpace(part), 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
out = append(out, id)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func joinPausedForwardIDs(ids []int64) string {
|
||||
if len(ids) == 0 {
|
||||
return ""
|
||||
}
|
||||
parts := make([]string, 0, len(ids))
|
||||
seen := make(map[int64]struct{}, len(ids))
|
||||
for _, id := range ids {
|
||||
if id <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[id]; ok {
|
||||
continue
|
||||
}
|
||||
seen[id] = struct{}{}
|
||||
parts = append(parts, strconv.FormatInt(id, 10))
|
||||
}
|
||||
return strings.Join(parts, ",")
|
||||
}
|
||||
|
||||
func (r *Repository) loadOrCreateTunnelQuotaTx(tx *gorm.DB, tunnelID int64, now time.Time) (*model.TunnelQuota, 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
|
||||
if err == nil {
|
||||
return q, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
nowMs := now.UnixMilli()
|
||||
q = &model.TunnelQuota{
|
||||
TunnelID: tunnelID,
|
||||
DayKey: dayKey,
|
||||
MonthKey: monthKey,
|
||||
CreatedTime: nowMs,
|
||||
UpdatedTime: nowMs,
|
||||
PausedForwardIDs: "",
|
||||
}
|
||||
if err := tx.Create(q).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return q, nil
|
||||
}
|
||||
|
||||
func applyTunnelQuotaWindowRoll(q *model.TunnelQuota, now time.Time) bool {
|
||||
if q == nil {
|
||||
return false
|
||||
}
|
||||
changed := false
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(now)
|
||||
if q.DayKey != dayKey {
|
||||
q.DayKey = dayKey
|
||||
q.DailyUsedBytes = 0
|
||||
changed = true
|
||||
}
|
||||
if q.MonthKey != monthKey {
|
||||
q.MonthKey = monthKey
|
||||
q.MonthlyUsedBytes = 0
|
||||
changed = true
|
||||
}
|
||||
return changed
|
||||
}
|
||||
|
||||
func (r *Repository) SaveTunnelQuotaConfigTx(tx *gorm.DB, tunnelID, 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 dailyLimitGB < 0 || monthlyLimitGB < 0 {
|
||||
return errors.New("quota limit cannot be negative")
|
||||
}
|
||||
current := time.UnixMilli(now)
|
||||
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, current)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
updates := map[string]interface{}{
|
||||
"daily_limit_gb": dailyLimitGB,
|
||||
"monthly_limit_gb": monthlyLimitGB,
|
||||
"updated_time": now,
|
||||
}
|
||||
if q.DayKey == 0 || q.MonthKey == 0 {
|
||||
dayKey, monthKey := tunnelQuotaWindowKeys(current)
|
||||
updates["day_key"] = dayKey
|
||||
updates["month_key"] = monthKey
|
||||
}
|
||||
return tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(updates).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ListTunnelQuotaViewsByTunnelIDs(tunnelIDs []int64, now time.Time) (map[int64]*model.TunnelQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
out := make(map[int64]*model.TunnelQuotaView)
|
||||
if len(tunnelIDs) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
var rows []model.TunnelQuota
|
||||
if err := r.db.Where("tunnel_id IN ?", tunnelIDs).Find(&rows).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, row := range rows {
|
||||
out[row.TunnelID] = normalizeTunnelQuotaView(cloneTunnelQuotaView(row), now)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetTunnelQuotaView(tunnelID int64, now time.Time) (*model.TunnelQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var row model.TunnelQuota
|
||||
err := r.db.Where("tunnel_id = ?", tunnelID).First(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeTunnelQuotaView(cloneTunnelQuotaView(row), now), nil
|
||||
}
|
||||
|
||||
func (r *Repository) AddTunnelQuotaUsage(tunnelID int64, usedBytes int64, now time.Time) (*model.TunnelQuotaView, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
if tunnelID <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
result := &model.TunnelQuotaView{}
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyTunnelQuotaWindowRoll(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{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
*result = *cloneTunnelQuotaView(*q)
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalizeTunnelQuotaView(result, now), nil
|
||||
}
|
||||
|
||||
func (r *Repository) MarkTunnelQuotaDisabled(tunnelID 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")
|
||||
}
|
||||
return r.db.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", tunnelID).Updates(map[string]interface{}{
|
||||
"disabled_by_quota": 1,
|
||||
"disabled_at": now,
|
||||
"paused_forward_ids": joinPausedForwardIDs(pausedForwardIDs),
|
||||
"updated_time": now,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) ResetTunnelQuotaUsage(tunnelID int64, scope string, now time.Time) (*TunnelQuotaRelease, 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")
|
||||
}
|
||||
scope = strings.TrimSpace(strings.ToLower(scope))
|
||||
if scope == "" {
|
||||
scope = "all"
|
||||
}
|
||||
if scope != "daily" && scope != "monthly" && scope != "all" {
|
||||
return nil, fmt.Errorf("unsupported quota reset scope: %s", scope)
|
||||
}
|
||||
var release *TunnelQuotaRelease
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
q, err := r.loadOrCreateTunnelQuotaTx(tx, tunnelID, now)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
applyTunnelQuotaWindowRoll(q, now)
|
||||
switch scope {
|
||||
case "daily":
|
||||
q.DailyUsedBytes = 0
|
||||
case "monthly":
|
||||
q.MonthlyUsedBytes = 0
|
||||
case "all":
|
||||
q.DailyUsedBytes = 0
|
||||
q.MonthlyUsedBytes = 0
|
||||
}
|
||||
q.UpdatedTime = now.UnixMilli()
|
||||
release = &TunnelQuotaRelease{TunnelID: tunnelID}
|
||||
if q.DisabledByQuota == 1 && !tunnelQuotaExceeded(cloneTunnelQuotaView(*q)) {
|
||||
release.EnableTunnel = 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{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"disabled_by_quota": q.DisabledByQuota,
|
||||
"disabled_at": q.DisabledAt,
|
||||
"paused_forward_ids": q.PausedForwardIDs,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return release, nil
|
||||
}
|
||||
|
||||
func (r *Repository) RollTunnelQuotaWindows(now time.Time) ([]TunnelQuotaRelease, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("repository not initialized")
|
||||
}
|
||||
var releases []TunnelQuotaRelease
|
||||
err := r.db.Transaction(func(tx *gorm.DB) error {
|
||||
var rows []model.TunnelQuota
|
||||
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
|
||||
release.ForwardIDs = parsePausedForwardIDs(q.PausedForwardIDs)
|
||||
q.DisabledByQuota = 0
|
||||
q.DisabledAt = 0
|
||||
q.PausedForwardIDs = ""
|
||||
changed = true
|
||||
}
|
||||
if !changed {
|
||||
continue
|
||||
}
|
||||
q.UpdatedTime = nowMs
|
||||
if err := tx.Model(&model.TunnelQuota{}).Where("tunnel_id = ?", q.TunnelID).Updates(map[string]interface{}{
|
||||
"daily_used_bytes": q.DailyUsedBytes,
|
||||
"monthly_used_bytes": q.MonthlyUsedBytes,
|
||||
"day_key": q.DayKey,
|
||||
"month_key": q.MonthKey,
|
||||
"disabled_by_quota": q.DisabledByQuota,
|
||||
"disabled_at": q.DisabledAt,
|
||||
"paused_forward_ids": q.PausedForwardIDs,
|
||||
"updated_time": q.UpdatedTime,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if release.EnableTunnel {
|
||||
releases = append(releases, release)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user