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
+38 -27
View File
@@ -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 ────────────────────────────────────────────────────
@@ -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)
}
}
@@ -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")
@@ -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)
}
}
@@ -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")