feat(nftables): ingest traffic counters

This commit is contained in:
sagitchu
2026-06-06 21:03:19 +08:00
committed by sagit
parent e569aedd3e
commit 006ea97200
10 changed files with 953 additions and 92 deletions
+18 -2
View File
@@ -20,7 +20,7 @@ func (h *Handler) StartBackgroundJobs() {
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
h.jobsCancel = cancel h.jobsCancel = cancel
h.jobsStarted = true h.jobsStarted = true
h.jobsWG.Add(7) h.jobsWG.Add(8)
h.jobsMu.Unlock() h.jobsMu.Unlock()
go h.runHourlyStatsLoop(ctx) go h.runHourlyStatsLoop(ctx)
@@ -30,6 +30,7 @@ func (h *Handler) StartBackgroundJobs() {
go h.runHealthChecks(ctx) go h.runHealthChecks(ctx)
go h.runTunnelQualityProber(ctx) go h.runTunnelQualityProber(ctx)
go h.runValidateLicenseJob(ctx) go h.runValidateLicenseJob(ctx)
go h.runNftablesTrafficCollectLoop(ctx)
} }
func (h *Handler) runValidateLicenseJob(ctx context.Context) { func (h *Handler) runValidateLicenseJob(ctx context.Context) {
@@ -64,7 +65,7 @@ func (h *Handler) validateLicenseJob() {
fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint") fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint")
client := license.NewKeygenClient(accountID, "") client := license.NewKeygenClient(accountID, "")
valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint) valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint)
if err != nil { if err != nil {
// Network error or timeout. Grace period by not revoking immediately here. // Network error or timeout. Grace period by not revoking immediately here.
return return
@@ -128,6 +129,21 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) {
h.qualityProber.Start(ctx) h.qualityProber.Start(ctx)
} }
func (h *Handler) runNftablesTrafficCollectLoop(ctx context.Context) {
defer h.jobsWG.Done()
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
h.runNftablesTrafficCollectJob(time.Now())
}
}
}
func (h *Handler) runHourlyStatsLoop(ctx context.Context) { func (h *Handler) runHourlyStatsLoop(ctx context.Context) {
defer h.jobsWG.Done() defer h.jobsWG.Done()
@@ -21,6 +21,7 @@ type nftablesRuntimeManager interface {
Test(ctx context.Context, cfg runtimenft.SSHConfig) error Test(ctx context.Context, cfg runtimenft.SSHConfig) error
Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error)
Clear(ctx context.Context, cfg runtimenft.SSHConfig) error Clear(ctx context.Context, cfg runtimenft.SSHConfig) error
CollectCounters(ctx context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error)
} }
func isNftablesForwardMode(mode string) bool { func isNftablesForwardMode(mode string) bool {
@@ -20,13 +20,16 @@ import (
) )
type fakeNftablesManager struct { type fakeNftablesManager struct {
testErr error testErr error
reconcileErr error reconcileErr error
reconcileHit int reconcileHit int
clearErr error clearErr error
clearHit int clearHit int
lastConfig runtimenft.SSHConfig collectErr error
lastPlan runtimenft.NodePlan collectHit int
counterSamples []runtimenft.CounterSample
lastConfig runtimenft.SSHConfig
lastPlan runtimenft.NodePlan
} }
func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error { func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error {
@@ -53,6 +56,15 @@ func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error
return f.clearErr return f.clearErr
} }
func (f *fakeNftablesManager) CollectCounters(_ context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) {
f.collectHit++
f.lastConfig = cfg
if f.collectErr != nil {
return nil, f.collectErr
}
return f.counterSamples, nil
}
type nftablesTestFixture struct { type nftablesTestFixture struct {
handler *Handler handler *Handler
nodeID int64 nodeID int64
@@ -1,9 +1,12 @@
package handler package handler
import ( import (
"context"
"log"
"math" "math"
"sort" "sort"
"strings" "strings"
"time"
runtimenft "go-backend/internal/runtime/nftables" runtimenft "go-backend/internal/runtime/nftables"
"go-backend/internal/store/model" "go-backend/internal/store/model"
@@ -22,6 +25,182 @@ type nftCounterStateKey struct {
direction string direction string
} }
func (h *Handler) runNftablesTrafficCollectJob(now time.Time) {
if h == nil || h.repo == nil {
return
}
nodes, err := h.repo.ListNftablesNodesForCollection()
if err != nil {
log.Printf("nftables traffic collection failed op=list_nodes err=%v", err)
return
}
for i := range nodes {
node := &nodes[i]
h.collectNftablesNodeTraffic(node.NodeID, &node.Config, now)
}
}
func (h *Handler) collectNftablesNodeTraffic(nodeID int64, cfgModel *model.NodeSSHConfig, now time.Time) {
if h == nil || h.repo == nil {
return
}
if h.nftablesManager == nil {
log.Printf("nftables traffic collection failed op=collect node_id=%d err=%v", nodeID, "nftables manager not initialized")
return
}
sshCfg, err := sshConfigFromModel(cfgModel)
if err != nil {
log.Printf("nftables traffic collection failed op=ssh_config node_id=%d err=%v", nodeID, err)
return
}
samples, err := h.nftablesManager.CollectCounters(context.Background(), sshCfg)
if err != nil {
log.Printf("nftables traffic collection failed op=collect node_id=%d err=%v", nodeID, err)
return
}
oldStates, err := h.repo.GetNftCounterStatesByNode(nodeID)
if err != nil {
log.Printf("nftables traffic collection failed op=list_states node_id=%d err=%v", nodeID, err)
return
}
bindings, err := h.repo.ListNftRuleBindingsByNode(nodeID)
if err != nil {
log.Printf("nftables traffic collection failed op=list_bindings node_id=%d err=%v", nodeID, err)
return
}
hashes := make(map[int64]string, len(bindings))
for _, binding := range bindings {
if strings.ToLower(strings.TrimSpace(binding.Status)) != runtimenft.StatusApplied {
continue
}
ruleHash := strings.TrimSpace(binding.RuleHash)
if ruleHash == "" {
continue
}
hashes[binding.ForwardID] = ruleHash
}
nowMs := now.UnixMilli()
boundSamples := filterNftCounterSamplesWithBinding(samples, hashes)
deltas, newStates := buildNftCounterDeltas(nodeID, boundSamples, oldStates, hashes, nowMs)
if len(newStates) == 0 {
if len(deltas) != 0 {
log.Printf("nftables traffic collection skipped suspicious deltas without states node_id=%d deltas=%d", nodeID, len(deltas))
}
return
}
var metas map[int64]repo.FlowUploadForwardMeta
forwardIDs := make([]int64, 0, len(deltas))
for _, delta := range deltas {
if delta.ForwardID > 0 {
forwardIDs = append(forwardIDs, delta.ForwardID)
}
}
if len(deltas) != 0 {
metas, err = h.repo.GetFlowUploadForwardMetas(forwardIDs)
if err != nil {
log.Printf("nftables traffic collection failed op=load_flow_metas node_id=%d err=%v", nodeID, err)
return
}
if missingForwardID, ok := firstNftDeltaMissingMeta(deltas, metas); ok {
log.Printf("nftables traffic collection skipped state advance op=missing_flow_meta node_id=%d forward_id=%d", nodeID, missingForwardID)
return
}
}
if len(deltas) == 0 {
if err := h.repo.UpsertNftCounterStates(newStates, nowMs); err != nil {
log.Printf("nftables traffic collection failed op=upsert_states node_id=%d err=%v", nodeID, err)
return
}
return
}
batch := buildNftFlowUploadBatch(deltas, metas)
if missingForwardID, ok := firstNftBatchMissingDelta(deltas, batch); ok {
log.Printf("nftables traffic collection skipped state advance op=unaccounted_delta node_id=%d forward_id=%d", nodeID, missingForwardID)
return
}
quotaViews, err := h.repo.ApplyNftTrafficAccounting(batch.flowDeltas, batch.quotaUsage, newStates, now)
if err != nil {
log.Printf("nftables traffic collection failed op=accounting node_id=%d err=%v", nodeID, err)
return
}
h.recordTunnelMetricsFromForwardBatch(nodeID, batch.forwardTraffic, metas, nowMs)
for userID, quota := range quotaViews {
h.enforceUserQuotaIfNeeded(userID, quota)
}
for _, target := range batch.policyTargets {
if target.UserID <= 0 || target.UserTunnelID <= 0 {
continue
}
h.enforceFlowPolicies(target.UserID, target.UserTunnelID)
}
}
func firstNftBatchMissingDelta(deltas []nftTrafficDelta, batch flowUploadBatch) (int64, bool) {
flowSeen := make(map[int64]struct{}, len(batch.flowDeltas))
for _, delta := range batch.flowDeltas {
flowSeen[delta.ForwardID] = struct{}{}
}
expectedRaw := make(map[int64]tunnelTrafficDelta, len(batch.forwardTraffic))
for _, delta := range deltas {
if delta.ForwardID <= 0 || (delta.BytesIn == 0 && delta.BytesOut == 0) {
continue
}
if delta.BytesIn < 0 || delta.BytesOut < 0 {
return delta.ForwardID, true
}
raw := expectedRaw[delta.ForwardID]
if raw.bytesIn > math.MaxInt64-delta.BytesIn || raw.bytesOut > math.MaxInt64-delta.BytesOut {
return delta.ForwardID, true
}
raw.bytesIn += delta.BytesIn
raw.bytesOut += delta.BytesOut
expectedRaw[delta.ForwardID] = raw
}
for forwardID, expected := range expectedRaw {
actual, ok := batch.forwardTraffic[forwardID]
if !ok || actual.bytesIn != expected.bytesIn || actual.bytesOut != expected.bytesOut {
return forwardID, true
}
if expected.bytesIn != 0 || expected.bytesOut != 0 {
if _, ok := flowSeen[forwardID]; !ok {
return forwardID, true
}
}
}
return 0, false
}
func firstNftDeltaMissingMeta(deltas []nftTrafficDelta, metas map[int64]repo.FlowUploadForwardMeta) (int64, bool) {
for _, delta := range deltas {
if delta.ForwardID <= 0 {
continue
}
if _, ok := metas[delta.ForwardID]; !ok {
return delta.ForwardID, true
}
}
return 0, false
}
func filterNftCounterSamplesWithBinding(samples []runtimenft.CounterSample, hashes map[int64]string) []runtimenft.CounterSample {
if len(samples) == 0 || len(hashes) == 0 {
return nil
}
filtered := make([]runtimenft.CounterSample, 0, len(samples))
for _, sample := range samples {
if _, ok := hashes[sample.ForwardID]; !ok {
continue
}
filtered = append(filtered, sample)
}
return filtered
}
func nftCounterKey(forwardID int64, protocol, direction string) nftCounterStateKey { func nftCounterKey(forwardID int64, protocol, direction string) nftCounterStateKey {
return nftCounterStateKey{ return nftCounterStateKey{
forwardID: forwardID, forwardID: forwardID,
@@ -1,8 +1,10 @@
package handler package handler
import ( import (
"errors"
"math" "math"
"testing" "testing"
"time"
runtimenft "go-backend/internal/runtime/nftables" runtimenft "go-backend/internal/runtime/nftables"
"go-backend/internal/store/model" "go-backend/internal/store/model"
@@ -312,3 +314,367 @@ func TestBuildNftFlowUploadBatchSkipsMissingMeta(t *testing.T) {
t.Fatalf("expected missing meta forward to be skipped from raw traffic") t.Fatalf("expected missing meta forward to be skipped from raw traffic")
} }
} }
func TestNftBatchCoversDeltasRequiresRawAndFlowEntries(t *testing.T) {
deltas := []nftTrafficDelta{{ForwardID: 20, BytesIn: 1, BytesOut: 0}}
batch := flowUploadBatch{
forwardTraffic: map[int64]tunnelTrafficDelta{20: {bytesIn: 1}},
flowDeltas: []repo.FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 1}},
}
if missing, ok := firstNftBatchMissingDelta(deltas, batch); ok || missing != 0 {
t.Fatalf("expected batch to cover delta, missing=%d ok=%v", missing, ok)
}
delete(batch.forwardTraffic, 20)
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
t.Fatalf("expected missing raw traffic for forward 20, got missing=%d ok=%v", missing, ok)
}
batch.forwardTraffic[20] = tunnelTrafficDelta{bytesIn: 1}
batch.flowDeltas = nil
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
t.Fatalf("expected missing flow delta for forward 20, got missing=%d ok=%v", missing, ok)
}
}
func TestNftBatchCoversDeltasRequiresAggregateRawTotals(t *testing.T) {
deltas := []nftTrafficDelta{
{ForwardID: 20, BytesIn: math.MaxInt64, BytesOut: 0},
{ForwardID: 20, BytesIn: 1, BytesOut: 0},
}
batch := buildNftFlowUploadBatch(deltas, map[int64]repo.FlowUploadForwardMeta{
20: {ForwardID: 20, UserID: 2, UserTunnelID: 10, TrafficRatio: 0.5, TunnelFlow: 1},
})
if missing, ok := firstNftBatchMissingDelta(deltas, batch); !ok || missing != 20 {
t.Fatalf("expected aggregate raw overflow/mismatch for forward 20, got missing=%d ok=%v", missing, ok)
}
}
func TestCollectNftablesNodeTrafficFirstBaselineSavesStateWithoutFlow(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
manager := &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2000, Packets: 20},
}}
h.nftablesManager = manager
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
if manager.collectHit != 1 {
t.Fatalf("expected one collection, got %d", manager.collectHit)
}
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 2 {
t.Fatalf("expected two baseline states, got %+v", states)
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
t.Fatalf("expected no forward flow on baseline, got %d", got)
}
if got := mustHandlerCount(t, h, `SELECT out_flow FROM user WHERE id = 1`); got != 0 {
t.Fatalf("expected no user flow on baseline, got %d", got)
}
}
func TestCollectNftablesNodeTrafficGrowthAppliesFlowAndUpdatesState(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
manager := &fakeNftablesManager{}
h.nftablesManager = manager
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
manager.counterSamples = []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2000, Packets: 20},
}
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
manager.counterSamples = []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionFromTarget, Bytes: 2600, Packets: 26},
}
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 400 {
t.Fatalf("expected forward in_flow=400, got %d", got)
}
if got := mustHandlerCount(t, h, `SELECT out_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 600 {
t.Fatalf("expected forward out_flow=600, got %d", got)
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM user WHERE id = 1`); got != 400 {
t.Fatalf("expected user in_flow=400, got %d", got)
}
if got := mustHandlerCount(t, h, `SELECT out_flow FROM user_tunnel WHERE id = ?`, fixture.userTunnelID); got != 600 {
t.Fatalf("expected user_tunnel out_flow=600, got %d", got)
}
if got := mustHandlerCount(t, h, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 1), 0)`); got != 1000 {
t.Fatalf("expected daily quota usage=1000, got %d", got)
}
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 2 {
t.Fatalf("expected two states after growth, got %+v", states)
}
for _, state := range states {
if state.Direction == runtimenft.CounterDirectionToTarget && state.Bytes != 1400 {
t.Fatalf("expected to-target state bytes 1400, got %+v", state)
}
if state.Direction == runtimenft.CounterDirectionFromTarget && state.Bytes != 2600 {
t.Fatalf("expected from-target state bytes 2600, got %+v", state)
}
}
}
func TestCollectNftablesNodeTrafficSkippedBatchDeltaDoesNotAdvanceState(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
if err := h.repo.DB().Exec(`UPDATE tunnel SET traffic_ratio = 2 WHERE id = (SELECT tunnel_id FROM forward WHERE id = ?)`, fixture.forwardID).Error; err != nil {
t.Fatalf("update tunnel ratio: %v", err)
}
manager := &fakeNftablesManager{}
h.nftablesManager = manager
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
manager.counterSamples = []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 0, Packets: 0},
}
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
manager.counterSamples = []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: uint64(math.MaxInt64), Packets: 1},
}
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 1 {
t.Fatalf("expected one state, got %+v", states)
}
if states[0].Bytes != 0 || states[0].Packets != 0 {
t.Fatalf("expected state to remain at old baseline after skipped batch delta, got %+v", states[0])
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
t.Fatalf("expected no forward flow for skipped batch delta, got %d", got)
}
}
func TestCollectNftablesNodeTrafficMetadataErrorDoesNotAdvanceState(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
manager := &fakeNftablesManager{}
h.nftablesManager = manager
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
manager.counterSamples = []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
}
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
if err := h.repo.DB().Exec(`DROP TABLE tunnel`).Error; err != nil {
t.Fatalf("drop tunnel table: %v", err)
}
manager.counterSamples = []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
}
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 1 {
t.Fatalf("expected one baseline state, got %+v", states)
}
if states[0].Bytes != 1000 || states[0].Packets != 10 {
t.Fatalf("expected state to remain at first baseline after metadata failure, got %+v", states[0])
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
t.Fatalf("expected no flow after metadata failure, got %d", got)
}
}
func TestCollectNftablesNodeTrafficMissingMetaDoesNotAdvanceState(t *testing.T) {
fixture := setupNftablesHandler(t)
h := fixture.handler
seedNftablesSSHConfig(t, h, fixture.nodeID)
forwardID := int64(4242)
nowMs := time.Now().UnixMilli()
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
ForwardID: forwardID,
NodeID: fixture.nodeID,
InPort: 20000,
Protocols: "tcp",
TargetAddr: "203.0.113.9:8080",
RuleHash: "hash-a",
Status: runtimenft.StatusApplied,
}, nowMs); err != nil {
t.Fatalf("seed stale applied binding: %v", err)
}
if err := h.repo.UpsertNftCounterStates([]repo.NftCounterStateInput{{
NodeID: fixture.nodeID,
ForwardID: forwardID,
Protocol: "tcp",
Direction: runtimenft.CounterDirectionToTarget,
RuleHash: "hash-a",
Bytes: 1000,
Packets: 10,
CollectedTime: nowMs,
}}, nowMs); err != nil {
t.Fatalf("seed counter state: %v", err)
}
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
{ForwardID: forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1400, Packets: 14},
}}
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000060, 0))
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 1 {
t.Fatalf("expected one state, got %+v", states)
}
if states[0].Bytes != 1000 || states[0].Packets != 10 {
t.Fatalf("expected state to remain at old baseline when meta is missing, got %+v", states[0])
}
}
func TestCollectNftablesNodeTrafficSkipsSamplesWithoutBinding(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
if err := h.repo.DeleteNftRuleBindingsByForward(fixture.forwardID); err != nil {
t.Fatalf("delete nft binding: %v", err)
}
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
}}
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 0 {
t.Fatalf("expected no state for unbound sample, got %+v", states)
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
t.Fatalf("expected no flow for unbound sample, got %d", got)
}
}
func TestCollectNftablesNodeTrafficSkipsNonAppliedBinding(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
if err := h.repo.MarkNftRuleBindingError(fixture.forwardID, fixture.nodeID, "apply failed", time.Now().UnixMilli()); err != nil {
t.Fatalf("mark binding error: %v", err)
}
h.nftablesManager = &fakeNftablesManager{counterSamples: []runtimenft.CounterSample{
{ForwardID: fixture.forwardID, Protocol: "tcp", Direction: runtimenft.CounterDirectionToTarget, Bytes: 1000, Packets: 10},
}}
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 0 {
t.Fatalf("expected no state for non-applied binding, got %+v", states)
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
t.Fatalf("expected no flow for non-applied binding, got %d", got)
}
}
func TestCollectNftablesNodeTrafficCollectionErrorDoesNotWriteState(t *testing.T) {
fixture := setupNftablesCollectionFixture(t)
h := fixture.handler
h.nftablesManager = &fakeNftablesManager{collectErr: errors.New("ssh failed")}
cfg := mustCollectionSSHConfig(t, h, fixture.nodeID)
h.collectNftablesNodeTraffic(fixture.nodeID, cfg, time.Unix(1700000000, 0))
states, err := h.repo.GetNftCounterStatesByNode(fixture.nodeID)
if err != nil {
t.Fatalf("load states: %v", err)
}
if len(states) != 0 {
t.Fatalf("expected no state on collection error, got %+v", states)
}
if got := mustHandlerCount(t, h, `SELECT in_flow FROM forward WHERE id = ?`, fixture.forwardID); got != 0 {
t.Fatalf("expected no flow on collection error, got %d", got)
}
}
type nftablesCollectionFixture struct {
handler *Handler
nodeID int64
forwardID int64
userTunnelID int64
}
func setupNftablesCollectionFixture(t *testing.T) nftablesCollectionFixture {
t.Helper()
fixture := setupNftablesHandler(t)
h := fixture.handler
seedNftablesSSHConfig(t, h, fixture.nodeID)
tunnelID := seedTunnelForNftables(t, h, "nft-traffic-tunnel", fixture.nodeID)
now := time.Now().UnixMilli()
if err := h.repo.DB().Exec(`
INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
VALUES(1, ?, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)
`, tunnelID).Error; err != nil {
t.Fatalf("seed user_tunnel: %v", err)
}
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
ForwardID: forward.ID,
NodeID: fixture.nodeID,
InPort: 20000,
Protocols: "tcp,udp",
TargetAddr: "203.0.113.9:8080",
RuleHash: "hash-a",
Status: runtimenft.StatusApplied,
}, now); err != nil {
t.Fatalf("seed nft binding: %v", err)
}
userTunnelID := mustHandlerCount(t, h, `SELECT id FROM user_tunnel WHERE user_id = 1 AND tunnel_id = ?`, tunnelID)
return nftablesCollectionFixture{
handler: h,
nodeID: fixture.nodeID,
forwardID: forward.ID,
userTunnelID: userTunnelID,
}
}
func mustCollectionSSHConfig(t *testing.T, h *Handler, nodeID int64) *model.NodeSSHConfig {
t.Helper()
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
if err != nil {
t.Fatalf("load ssh config: %v", err)
}
return cfg
}
func mustHandlerCount(t *testing.T, h *Handler, query string, args ...interface{}) int64 {
t.Helper()
var value int64
if err := h.repo.DB().Raw(query, args...).Row().Scan(&value); err != nil {
t.Fatalf("query %q failed: %v", query, err)
}
return value
}
+38 -27
View File
@@ -104,6 +104,19 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
return nil return nil
} }
return r.db.Transaction(func(tx *gorm.DB) error {
return applyFlowUploadDeltasTx(tx, deltas)
})
}
func applyFlowUploadDeltasTx(tx *gorm.DB, deltas []FlowUploadCounterDelta) error {
if tx == nil {
return errors.New("database unavailable")
}
if len(deltas) == 0 {
return nil
}
forwardTotals := make(map[int64][2]int64, len(deltas)) forwardTotals := make(map[int64][2]int64, len(deltas))
userTotals := make(map[int64][2]int64, len(deltas)) userTotals := make(map[int64][2]int64, len(deltas))
userTunnelTotals := make(map[int64][2]int64, len(deltas)) userTunnelTotals := make(map[int64][2]int64, len(deltas))
@@ -128,36 +141,34 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta)
} }
} }
return r.db.Transaction(func(tx *gorm.DB) error { for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) {
for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) { total := forwardTotals[forwardID]
total := forwardTotals[forwardID] if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{
if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{ "in_flow": gorm.Expr("in_flow + ?", total[0]),
"in_flow": gorm.Expr("in_flow + ?", total[0]), "out_flow": gorm.Expr("out_flow + ?", total[1]),
"out_flow": gorm.Expr("out_flow + ?", total[1]), }).Error; err != nil {
}).Error; err != nil { return err
return err
}
} }
for _, userID := range sortedFlowUploadTargetIDs(userTotals) { }
total := userTotals[userID] for _, userID := range sortedFlowUploadTargetIDs(userTotals) {
if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{ total := userTotals[userID]
"in_flow": gorm.Expr("in_flow + ?", total[0]), if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{
"out_flow": gorm.Expr("out_flow + ?", total[1]), "in_flow": gorm.Expr("in_flow + ?", total[0]),
}).Error; err != nil { "out_flow": gorm.Expr("out_flow + ?", total[1]),
return err }).Error; err != nil {
} return err
} }
for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) { }
total := userTunnelTotals[userTunnelID] for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) {
if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{ total := userTunnelTotals[userTunnelID]
"in_flow": gorm.Expr("in_flow + ?", total[0]), if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{
"out_flow": gorm.Expr("out_flow + ?", total[1]), "in_flow": gorm.Expr("in_flow + ?", total[0]),
}).Error; err != nil { "out_flow": gorm.Expr("out_flow + ?", total[1]),
return err }).Error; err != nil {
} return err
} }
return nil }
}) return nil
} }
// ─── Open / Close ──────────────────────────────────────────────────── // ─── Open / Close ────────────────────────────────────────────────────
@@ -86,6 +86,99 @@ func TestGetFlowUploadForwardMetasAndApplyFlowUploadDeltasBatch(t *testing.T) {
} }
} }
func TestApplyNftTrafficAccountingAppliesFlowQuotaAndStates(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
seedFlowBatchRows(t, r, nowMs)
quotaViews, err := r.ApplyNftTrafficAccounting(
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
map[int64]int64{2: 1140},
[]NftCounterStateInput{{
NodeID: 11,
ForwardID: 20,
Protocol: "tcp",
Direction: "to-target",
RuleHash: "hash-a",
Bytes: 1400,
Packets: 14,
CollectedTime: nowMs,
}},
now,
)
if err != nil {
t.Fatalf("ApplyNftTrafficAccounting: %v", err)
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 480 {
t.Fatalf("expected forward in_flow=480, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 660 {
t.Fatalf("expected user out_flow=660, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM user_tunnel WHERE id = 10`); got != 480 {
t.Fatalf("expected user_tunnel in_flow=480, got %d", got)
}
if quotaViews[2] == nil || quotaViews[2].DailyUsedBytes != 1140 || quotaViews[2].MonthlyUsedBytes != 1140 {
t.Fatalf("unexpected quota view: %#v", quotaViews[2])
}
states, err := r.GetNftCounterStatesByNode(11)
if err != nil {
t.Fatalf("GetNftCounterStatesByNode: %v", err)
}
if len(states) != 1 || states[0].ForwardID != 20 || states[0].Bytes != 1400 {
t.Fatalf("unexpected nft counter state: %+v", states)
}
}
func TestApplyNftTrafficAccountingRollsBackFlowAndQuotaWhenStateWriteFails(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "nft-accounting-rollback.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now()
nowMs := now.UnixMilli()
seedFlowBatchRows(t, r, nowMs)
if err := r.DB().Exec(`DROP TABLE nft_counter_state`).Error; err != nil {
t.Fatalf("drop nft_counter_state: %v", err)
}
_, err = r.ApplyNftTrafficAccounting(
[]FlowUploadCounterDelta{{ForwardID: 20, UserID: 2, UserTunnelID: 10, InFlow: 480, OutFlow: 660}},
map[int64]int64{2: 1140},
[]NftCounterStateInput{{
NodeID: 11,
ForwardID: 20,
Protocol: "tcp",
Direction: "to-target",
RuleHash: "hash-a",
Bytes: 1400,
Packets: 14,
CollectedTime: nowMs,
}},
now,
)
if err == nil {
t.Fatalf("expected ApplyNftTrafficAccounting to fail")
}
if got := mustFlowBatchCount(t, r, `SELECT in_flow FROM forward WHERE id = 20`); got != 0 {
t.Fatalf("expected forward flow rollback, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT out_flow FROM user WHERE id = 2`); got != 0 {
t.Fatalf("expected user flow rollback, got %d", got)
}
if got := mustFlowBatchCount(t, r, `SELECT COALESCE((SELECT daily_used_bytes FROM user_quota WHERE user_id = 2), 0)`); got != 0 {
t.Fatalf("expected quota rollback, got %d", got)
}
}
func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) { func TestGetFlowUploadForwardMetasKeepsForwardsWhenTunnelRowMissing(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db")) r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db"))
if err != nil { if err != nil {
@@ -167,3 +260,19 @@ func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...inter
} }
return value return value
} }
func seedFlowBatchRows(t *testing.T, r *Repository, now int64) {
t.Helper()
if err := r.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, 'u2', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)`, now, now).Error; err != nil {
t.Fatalf("insert user: %v", err)
}
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, 't1', 2.0, 1, 'tls', 3, ?, ?, 1, NULL, 0)`, now, now).Error; err != nil {
t.Fatalf("insert tunnel: %v", err)
}
if err := r.DB().Exec(`INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(10, 2, 1, NULL, 99999, 99999, 0, 0, 1, 2727251700000, 1)`).Error; err != nil {
t.Fatalf("insert user_tunnel: %v", err)
}
if err := r.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(20, 2, 'u2', 'f20', 1, '1.1.1.1:80', 'fifo', 0, 0, ?, ?, 1, 0)`, now, now).Error; err != nil {
t.Fatalf("insert forward: %v", err)
}
}
@@ -4,6 +4,7 @@ import (
"errors" "errors"
"math" "math"
"strings" "strings"
"time"
"go-backend/internal/store/model" "go-backend/internal/store/model"
@@ -30,6 +31,64 @@ type NftCounterStateInput struct {
CollectedTime int64 CollectedTime int64
} }
type NftablesCollectionNode struct {
NodeID int64
Config model.NodeSSHConfig
}
func (r *Repository) ListNftablesNodesForCollection() ([]NftablesCollectionNode, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
type collectionRow struct {
NodeID int64 `gorm:"column:node_id"`
ConfigID int64 `gorm:"column:config_id"`
Host string `gorm:"column:host"`
Port int `gorm:"column:port"`
Username string `gorm:"column:username"`
AuthType string `gorm:"column:auth_type"`
Password string `gorm:"column:password"`
PrivateKey string `gorm:"column:private_key"`
Passphrase string `gorm:"column:passphrase"`
SudoMode string `gorm:"column:sudo_mode"`
CreatedTime int64 `gorm:"column:created_time"`
UpdatedTime int64 `gorm:"column:updated_time"`
}
var rows []collectionRow
if err := r.db.Table("node").
Select("node.id AS node_id, node_ssh_config.id AS config_id, node_ssh_config.host, node_ssh_config.port, node_ssh_config.username, node_ssh_config.auth_type, node_ssh_config.password, node_ssh_config.private_key, node_ssh_config.passphrase, node_ssh_config.sudo_mode, node_ssh_config.created_time, node_ssh_config.updated_time").
Joins("JOIN node_ssh_config ON node_ssh_config.node_id = node.id").
Where("node.status = ? AND LOWER(TRIM(node.forward_mode)) = ?", 1, "nftables").
Order("node.id ASC").
Scan(&rows).Error; err != nil {
return nil, err
}
nodes := make([]NftablesCollectionNode, 0, len(rows))
for _, row := range rows {
nodes = append(nodes, NftablesCollectionNode{
NodeID: row.NodeID,
Config: model.NodeSSHConfig{
ID: row.ConfigID,
NodeID: row.NodeID,
Host: row.Host,
Port: row.Port,
Username: row.Username,
AuthType: row.AuthType,
Password: nullStringFromInterface(row.Password),
PrivateKey: nullStringFromInterface(row.PrivateKey),
Passphrase: nullStringFromInterface(row.Passphrase),
SudoMode: row.SudoMode,
CreatedTime: row.CreatedTime,
UpdatedTime: row.UpdatedTime,
},
})
}
return nodes, nil
}
func (r *Repository) GetNftCounterStatesByNode(nodeID int64) ([]model.NftCounterState, error) { func (r *Repository) GetNftCounterStatesByNode(nodeID int64) ([]model.NftCounterState, error) {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return nil, errors.New("repository not initialized") return nil, errors.New("repository not initialized")
@@ -50,33 +109,63 @@ func (r *Repository) UpsertNftCounterStates(inputs []NftCounterStateInput, now i
} }
return r.db.Transaction(func(tx *gorm.DB) error { return r.db.Transaction(func(tx *gorm.DB) error {
for _, input := range inputs { return upsertNftCounterStatesTx(tx, inputs, now)
row, ok := nftCounterStateFromInput(input, now)
if !ok {
continue
}
if err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{
{Name: "node_id"},
{Name: "forward_id"},
{Name: "protocol"},
{Name: "direction"},
},
DoUpdates: clause.Assignments(map[string]interface{}{
"rule_hash": row.RuleHash,
"bytes": row.Bytes,
"packets": row.Packets,
"collected_time": row.CollectedTime,
"updated_time": row.UpdatedTime,
}),
}).Create(&row).Error; err != nil {
return err
}
}
return nil
}) })
} }
func (r *Repository) ApplyNftTrafficAccounting(deltas []FlowUploadCounterDelta, quotaUsage map[int64]int64, states []NftCounterStateInput, now time.Time) (map[int64]*model.UserQuotaView, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
quotaViews := map[int64]*model.UserQuotaView{}
err := r.db.Transaction(func(tx *gorm.DB) error {
if err := applyFlowUploadDeltasTx(tx, deltas); err != nil {
return err
}
var err error
quotaViews, err = r.addUserQuotaUsageBatchTx(tx, quotaUsage, now)
if err != nil {
return err
}
return upsertNftCounterStatesTx(tx, states, now.UnixMilli())
})
if err != nil {
return nil, err
}
return quotaViews, nil
}
func upsertNftCounterStatesTx(tx *gorm.DB, inputs []NftCounterStateInput, now int64) error {
if tx == nil {
return errors.New("database unavailable")
}
for _, input := range inputs {
row, ok := nftCounterStateFromInput(input, now)
if !ok {
continue
}
if err := tx.Clauses(clause.OnConflict{
Columns: []clause.Column{
{Name: "node_id"},
{Name: "forward_id"},
{Name: "protocol"},
{Name: "direction"},
},
DoUpdates: clause.Assignments(map[string]interface{}{
"rule_hash": row.RuleHash,
"bytes": row.Bytes,
"packets": row.Packets,
"collected_time": row.CollectedTime,
"updated_time": row.UpdatedTime,
}),
}).Create(&row).Error; err != nil {
return err
}
}
return nil
}
func (r *Repository) DeleteNftCounterStatesByForward(forwardID int64) error { func (r *Repository) DeleteNftCounterStatesByForward(forwardID int64) error {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return errors.New("repository not initialized") return errors.New("repository not initialized")
@@ -2,7 +2,9 @@ package repo
import ( import (
"math" "math"
"path/filepath"
"testing" "testing"
"time"
) )
func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) { func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) {
@@ -159,3 +161,65 @@ func TestNftCounterStateUpsertSkipsCountersAboveInt64(t *testing.T) {
t.Fatalf("unexpected valid counter state row: %+v", rows[0]) t.Fatalf("unexpected valid counter state row: %+v", rows[0])
} }
} }
func TestListNftablesNodesForCollectionReturnsActiveNftablesWithSSHOrdered(t *testing.T) {
r, err := Open(filepath.Join(t.TempDir(), "nft-collection.db"))
if err != nil {
t.Fatalf("open repo: %v", err)
}
defer r.Close()
now := time.Now().UnixMilli()
seedCollectionNode(t, r, 1, "agent", 1, now)
seedCollectionNode(t, r, 2, " nftables ", 1, now)
seedCollectionNode(t, r, 3, "NFTABLES", 0, now)
seedCollectionNode(t, r, 4, "nftables", 1, now)
seedCollectionNode(t, r, 5, "nftables", 1, now)
if err := r.UpsertNodeSSHConfig(4, NftSSHConfigInput{
Host: "203.0.113.4",
Port: 2222,
Username: "root",
AuthType: "password",
Password: "secret-4",
SudoMode: "none",
}, now); err != nil {
t.Fatalf("upsert ssh config 4: %v", err)
}
if err := r.UpsertNodeSSHConfig(2, NftSSHConfigInput{
Host: "203.0.113.2",
Port: 22,
Username: "admin",
AuthType: "private_key",
SudoMode: "sudo",
}, now); err != nil {
t.Fatalf("upsert ssh config 2: %v", err)
}
nodes, err := r.ListNftablesNodesForCollection()
if err != nil {
t.Fatalf("ListNftablesNodesForCollection: %v", err)
}
if len(nodes) != 2 {
t.Fatalf("expected 2 collection nodes, got %d: %+v", len(nodes), nodes)
}
if nodes[0].NodeID != 2 || nodes[1].NodeID != 4 {
t.Fatalf("expected nodes ordered by id [2 4], got [%d %d]", nodes[0].NodeID, nodes[1].NodeID)
}
if nodes[0].Config.NodeID != 2 || nodes[0].Config.Host != "203.0.113.2" || nodes[0].Config.Username != "admin" {
t.Fatalf("unexpected first config: %+v", nodes[0].Config)
}
if nodes[1].Config.NodeID != 4 || nodes[1].Config.Port != 2222 || nodes[1].Config.Password.String != "secret-4" {
t.Fatalf("unexpected second config: %+v", nodes[1].Config)
}
}
func seedCollectionNode(t *testing.T, r *Repository, id int64, forwardMode string, status int, now int64) {
t.Helper()
if err := r.DB().Exec(`
INSERT INTO node(id, name, secret, server_ip, port, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx, forward_mode)
VALUES(?, ?, 'secret', ?, '1000-2000', ?, ?, ?, '[::]', '[::]', 0, ?)
`, id, "node", "198.51.100.1", now, now, status, forwardMode).Error; err != nil {
t.Fatalf("insert node %d: %v", id, err)
}
}
@@ -264,39 +264,11 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
return map[int64]*model.UserQuotaView{}, nil return map[int64]*model.UserQuotaView{}, nil
} }
result := make(map[int64]*model.UserQuotaView, len(usages)) var result map[int64]*model.UserQuotaView
err := r.db.Transaction(func(tx *gorm.DB) error { err := r.db.Transaction(func(tx *gorm.DB) error {
userIDs := make([]int64, 0, len(usages)) var err error
for userID := range usages { result, err = r.addUserQuotaUsageBatchTx(tx, usages, now)
if userID > 0 { return err
userIDs = append(userIDs, userID)
}
}
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
for _, userID := range userIDs {
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
if err != nil {
return err
}
applyUserQuotaWindowRoll(q, now)
if usages[userID] > 0 {
q.DailyUsedBytes += usages[userID]
q.MonthlyUsedBytes += usages[userID]
}
q.UpdatedTime = now.UnixMilli()
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,
"month_key": q.MonthKey,
"updated_time": q.UpdatedTime,
}).Error; err != nil {
return err
}
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
}
return nil
}) })
if err != nil { if err != nil {
return nil, err return nil, err
@@ -304,6 +276,48 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim
return result, nil return result, nil
} }
func (r *Repository) addUserQuotaUsageBatchTx(tx *gorm.DB, usages map[int64]int64, now time.Time) (map[int64]*model.UserQuotaView, error) {
if tx == nil {
return nil, errors.New("database unavailable")
}
if len(usages) == 0 {
return map[int64]*model.UserQuotaView{}, nil
}
result := make(map[int64]*model.UserQuotaView, len(usages))
userIDs := make([]int64, 0, len(usages))
for userID := range usages {
if userID > 0 {
userIDs = append(userIDs, userID)
}
}
sort.Slice(userIDs, func(i, j int) bool { return userIDs[i] < userIDs[j] })
for _, userID := range userIDs {
q, err := r.loadOrCreateUserQuotaTx(tx, userID, now)
if err != nil {
return nil, err
}
applyUserQuotaWindowRoll(q, now)
if usages[userID] > 0 {
q.DailyUsedBytes += usages[userID]
q.MonthlyUsedBytes += usages[userID]
}
q.UpdatedTime = now.UnixMilli()
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,
"month_key": q.MonthKey,
"updated_time": q.UpdatedTime,
}).Error; err != nil {
return nil, err
}
result[userID] = normalizeUserQuotaView(cloneUserQuotaView(*q), now)
}
return result, nil
}
func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error { func (r *Repository) MarkUserQuotaDisabled(userID int64, pausedForwardIDs []int64, now int64) error {
if r == nil || r.db == nil { if r == nil || r.db == nil {
return errors.New("repository not initialized") return errors.New("repository not initialized")