From 006ea9720014a67dcfb3b3c90f2ebbd8d7ce9b33 Mon Sep 17 00:00:00 2001 From: sagitchu <601096721@qq.com> Date: Sat, 6 Jun 2026 21:03:19 +0800 Subject: [PATCH] feat(nftables): ingest traffic counters --- go-backend/internal/http/handler/jobs.go | 20 +- .../internal/http/handler/nftables_runtime.go | 1 + .../http/handler/nftables_runtime_test.go | 26 +- .../internal/http/handler/nftables_traffic.go | 179 +++++++++ .../http/handler/nftables_traffic_test.go | 366 ++++++++++++++++++ go-backend/internal/store/repo/repository.go | 65 ++-- .../store/repo/repository_flow_batch_test.go | 109 ++++++ .../store/repo/repository_nft_counter.go | 137 +++++-- .../store/repo/repository_nft_counter_test.go | 64 +++ .../store/repo/repository_user_quota.go | 78 ++-- 10 files changed, 953 insertions(+), 92 deletions(-) diff --git a/go-backend/internal/http/handler/jobs.go b/go-backend/internal/http/handler/jobs.go index 8123c0b..c89f9e8 100644 --- a/go-backend/internal/http/handler/jobs.go +++ b/go-backend/internal/http/handler/jobs.go @@ -20,7 +20,7 @@ func (h *Handler) StartBackgroundJobs() { ctx, cancel := context.WithCancel(context.Background()) h.jobsCancel = cancel h.jobsStarted = true - h.jobsWG.Add(7) + h.jobsWG.Add(8) h.jobsMu.Unlock() go h.runHourlyStatsLoop(ctx) @@ -30,6 +30,7 @@ func (h *Handler) StartBackgroundJobs() { go h.runHealthChecks(ctx) go h.runTunnelQualityProber(ctx) go h.runValidateLicenseJob(ctx) + go h.runNftablesTrafficCollectLoop(ctx) } func (h *Handler) runValidateLicenseJob(ctx context.Context) { @@ -64,7 +65,7 @@ func (h *Handler) validateLicenseJob() { fingerprint, _ := h.repo.GetViteConfigValue("machine_fingerprint") client := license.NewKeygenClient(accountID, "") valResp, err := client.ValidateKeyWithFingerprint(key, fingerprint) - + if err != nil { // Network error or timeout. Grace period by not revoking immediately here. return @@ -128,6 +129,21 @@ func (h *Handler) runTunnelQualityProber(ctx context.Context) { 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) { defer h.jobsWG.Done() diff --git a/go-backend/internal/http/handler/nftables_runtime.go b/go-backend/internal/http/handler/nftables_runtime.go index fa70eb6..eed00d1 100644 --- a/go-backend/internal/http/handler/nftables_runtime.go +++ b/go-backend/internal/http/handler/nftables_runtime.go @@ -21,6 +21,7 @@ type nftablesRuntimeManager interface { Test(ctx context.Context, cfg runtimenft.SSHConfig) error Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) Clear(ctx context.Context, cfg runtimenft.SSHConfig) error + CollectCounters(ctx context.Context, cfg runtimenft.SSHConfig) ([]runtimenft.CounterSample, error) } func isNftablesForwardMode(mode string) bool { diff --git a/go-backend/internal/http/handler/nftables_runtime_test.go b/go-backend/internal/http/handler/nftables_runtime_test.go index b7a5d8c..14145f6 100644 --- a/go-backend/internal/http/handler/nftables_runtime_test.go +++ b/go-backend/internal/http/handler/nftables_runtime_test.go @@ -20,13 +20,16 @@ import ( ) type fakeNftablesManager struct { - testErr error - reconcileErr error - reconcileHit int - clearErr error - clearHit int - lastConfig runtimenft.SSHConfig - lastPlan runtimenft.NodePlan + testErr error + reconcileErr error + reconcileHit int + clearErr error + clearHit int + collectErr error + collectHit int + counterSamples []runtimenft.CounterSample + lastConfig runtimenft.SSHConfig + lastPlan runtimenft.NodePlan } 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 } +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 { handler *Handler nodeID int64 diff --git a/go-backend/internal/http/handler/nftables_traffic.go b/go-backend/internal/http/handler/nftables_traffic.go index e48ba02..4e2174e 100644 --- a/go-backend/internal/http/handler/nftables_traffic.go +++ b/go-backend/internal/http/handler/nftables_traffic.go @@ -1,9 +1,12 @@ package handler import ( + "context" + "log" "math" "sort" "strings" + "time" runtimenft "go-backend/internal/runtime/nftables" "go-backend/internal/store/model" @@ -22,6 +25,182 @@ type nftCounterStateKey struct { 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 { return nftCounterStateKey{ forwardID: forwardID, diff --git a/go-backend/internal/http/handler/nftables_traffic_test.go b/go-backend/internal/http/handler/nftables_traffic_test.go index e4c7cfc..abb532c 100644 --- a/go-backend/internal/http/handler/nftables_traffic_test.go +++ b/go-backend/internal/http/handler/nftables_traffic_test.go @@ -1,8 +1,10 @@ package handler import ( + "errors" "math" "testing" + "time" runtimenft "go-backend/internal/runtime/nftables" "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") } } + +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 +} diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index 46c9cc7..2629dac 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -104,6 +104,19 @@ func (r *Repository) ApplyFlowUploadDeltasBatch(deltas []FlowUploadCounterDelta) 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)) userTotals := 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) { - total := forwardTotals[forwardID] - if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{ - "in_flow": gorm.Expr("in_flow + ?", total[0]), - "out_flow": gorm.Expr("out_flow + ?", total[1]), - }).Error; err != nil { - return err - } + for _, forwardID := range sortedFlowUploadTargetIDs(forwardTotals) { + total := forwardTotals[forwardID] + if err := tx.Model(&model.Forward{}).Where("id = ?", forwardID).UpdateColumns(map[string]interface{}{ + "in_flow": gorm.Expr("in_flow + ?", total[0]), + "out_flow": gorm.Expr("out_flow + ?", total[1]), + }).Error; err != nil { + return err } - for _, userID := range sortedFlowUploadTargetIDs(userTotals) { - total := userTotals[userID] - if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{ - "in_flow": gorm.Expr("in_flow + ?", total[0]), - "out_flow": gorm.Expr("out_flow + ?", total[1]), - }).Error; err != nil { - return err - } + } + for _, userID := range sortedFlowUploadTargetIDs(userTotals) { + total := userTotals[userID] + if err := tx.Model(&model.User{}).Where("id = ?", userID).UpdateColumns(map[string]interface{}{ + "in_flow": gorm.Expr("in_flow + ?", total[0]), + "out_flow": gorm.Expr("out_flow + ?", total[1]), + }).Error; err != nil { + return err } - for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) { - total := userTunnelTotals[userTunnelID] - if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{ - "in_flow": gorm.Expr("in_flow + ?", total[0]), - "out_flow": gorm.Expr("out_flow + ?", total[1]), - }).Error; err != nil { - return err - } + } + for _, userTunnelID := range sortedFlowUploadTargetIDs(userTunnelTotals) { + total := userTunnelTotals[userTunnelID] + if err := tx.Model(&model.UserTunnel{}).Where("id = ?", userTunnelID).UpdateColumns(map[string]interface{}{ + "in_flow": gorm.Expr("in_flow + ?", total[0]), + "out_flow": gorm.Expr("out_flow + ?", total[1]), + }).Error; err != nil { + return err } - return nil - }) + } + return nil } // ─── Open / Close ──────────────────────────────────────────────────── diff --git a/go-backend/internal/store/repo/repository_flow_batch_test.go b/go-backend/internal/store/repo/repository_flow_batch_test.go index 928a729..8c1e5c5 100644 --- a/go-backend/internal/store/repo/repository_flow_batch_test.go +++ b/go-backend/internal/store/repo/repository_flow_batch_test.go @@ -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) { r, err := Open(filepath.Join(t.TempDir(), "flow-batch-missing-tunnel.db")) if err != nil { @@ -167,3 +260,19 @@ func mustFlowBatchCount(t *testing.T, r *Repository, query string, args ...inter } 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) + } +} diff --git a/go-backend/internal/store/repo/repository_nft_counter.go b/go-backend/internal/store/repo/repository_nft_counter.go index 66b30d4..e2530f5 100644 --- a/go-backend/internal/store/repo/repository_nft_counter.go +++ b/go-backend/internal/store/repo/repository_nft_counter.go @@ -4,6 +4,7 @@ import ( "errors" "math" "strings" + "time" "go-backend/internal/store/model" @@ -30,6 +31,64 @@ type NftCounterStateInput struct { 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) { if r == nil || r.db == nil { 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 { - 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 + return upsertNftCounterStatesTx(tx, inputs, now) }) } +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 { if r == nil || r.db == nil { return errors.New("repository not initialized") diff --git a/go-backend/internal/store/repo/repository_nft_counter_test.go b/go-backend/internal/store/repo/repository_nft_counter_test.go index d955843..02b9838 100644 --- a/go-backend/internal/store/repo/repository_nft_counter_test.go +++ b/go-backend/internal/store/repo/repository_nft_counter_test.go @@ -2,7 +2,9 @@ package repo import ( "math" + "path/filepath" "testing" + "time" ) func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) { @@ -159,3 +161,65 @@ func TestNftCounterStateUpsertSkipsCountersAboveInt64(t *testing.T) { 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) + } +} diff --git a/go-backend/internal/store/repo/repository_user_quota.go b/go-backend/internal/store/repo/repository_user_quota.go index 4c97c24..d117cf4 100644 --- a/go-backend/internal/store/repo/repository_user_quota.go +++ b/go-backend/internal/store/repo/repository_user_quota.go @@ -264,39 +264,11 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim 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 { - 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 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 + var err error + result, err = r.addUserQuotaUsageBatchTx(tx, usages, now) + return err }) if err != nil { return nil, err @@ -304,6 +276,48 @@ func (r *Repository) AddUserQuotaUsageBatch(usages map[int64]int64, now time.Tim 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 { if r == nil || r.db == nil { return errors.New("repository not initialized")