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())
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()
@@ -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 {
@@ -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
@@ -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,
@@ -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
}