feat(quota): add tunnel traffic quota with daily/monthly limits (#291) (#308)

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:
sagit
2026-03-11 16:09:03 +08:00
committed by GitHub
parent 69faeaa9a6
commit 5e96a8de72
16 changed files with 1133 additions and 13 deletions
@@ -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)
+1
View File
@@ -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
}