mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 18:56:37 +08:00
feat(nftables): persist counter state
This commit is contained in:
@@ -131,6 +131,22 @@ type NftRuleBinding struct {
|
|||||||
|
|
||||||
func (NftRuleBinding) TableName() string { return "nft_rule_binding" }
|
func (NftRuleBinding) TableName() string { return "nft_rule_binding" }
|
||||||
|
|
||||||
|
type NftCounterState struct {
|
||||||
|
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||||
|
NodeID int64 `gorm:"column:node_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||||
|
ForwardID int64 `gorm:"column:forward_id;not null;uniqueIndex:idx_nft_counter_state_key;index"`
|
||||||
|
Protocol string `gorm:"type:varchar(10);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||||
|
Direction string `gorm:"type:varchar(20);not null;uniqueIndex:idx_nft_counter_state_key"`
|
||||||
|
RuleHash string `gorm:"column:rule_hash;type:varchar(128);not null;default:''"`
|
||||||
|
Bytes int64 `gorm:"not null;default:0"`
|
||||||
|
Packets int64 `gorm:"not null;default:0"`
|
||||||
|
CollectedTime int64 `gorm:"column:collected_time;not null;default:0"`
|
||||||
|
CreatedTime int64 `gorm:"column:created_time;not null"`
|
||||||
|
UpdatedTime int64 `gorm:"column:updated_time;not null"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (NftCounterState) TableName() string { return "nft_counter_state" }
|
||||||
|
|
||||||
type SpeedLimit struct {
|
type SpeedLimit struct {
|
||||||
ID int64 `gorm:"primaryKey;autoIncrement"`
|
ID int64 `gorm:"primaryKey;autoIncrement"`
|
||||||
Name string `gorm:"type:varchar(100);not null"`
|
Name string `gorm:"type:varchar(100);not null"`
|
||||||
|
|||||||
@@ -282,6 +282,7 @@ func autoMigrateAll(db *gorm.DB) error {
|
|||||||
&model.Node{},
|
&model.Node{},
|
||||||
&model.NodeSSHConfig{},
|
&model.NodeSSHConfig{},
|
||||||
&model.NftRuleBinding{},
|
&model.NftRuleBinding{},
|
||||||
|
&model.NftCounterState{},
|
||||||
&model.SpeedLimit{},
|
&model.SpeedLimit{},
|
||||||
&model.StatisticsFlow{},
|
&model.StatisticsFlow{},
|
||||||
&model.Tunnel{},
|
&model.Tunnel{},
|
||||||
|
|||||||
@@ -0,0 +1,116 @@
|
|||||||
|
package repo
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"go-backend/internal/store/model"
|
||||||
|
|
||||||
|
"gorm.io/gorm"
|
||||||
|
"gorm.io/gorm/clause"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
nftCounterProtocolTCP = "tcp"
|
||||||
|
nftCounterProtocolUDP = "udp"
|
||||||
|
|
||||||
|
nftCounterDirectionToTarget = "to-target"
|
||||||
|
nftCounterDirectionFromTarget = "from-target"
|
||||||
|
)
|
||||||
|
|
||||||
|
type NftCounterStateInput struct {
|
||||||
|
NodeID int64
|
||||||
|
ForwardID int64
|
||||||
|
Protocol string
|
||||||
|
Direction string
|
||||||
|
RuleHash string
|
||||||
|
Bytes uint64
|
||||||
|
Packets uint64
|
||||||
|
CollectedTime int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Repository) GetNftCounterStatesByNode(nodeID int64) ([]model.NftCounterState, error) {
|
||||||
|
if r == nil || r.db == nil {
|
||||||
|
return nil, errors.New("repository not initialized")
|
||||||
|
}
|
||||||
|
var rows []model.NftCounterState
|
||||||
|
err := r.db.Where("node_id = ?", nodeID).
|
||||||
|
Order("forward_id ASC, protocol ASC, direction ASC").
|
||||||
|
Find(&rows).Error
|
||||||
|
return rows, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Repository) UpsertNftCounterStates(inputs []NftCounterStateInput, now int64) error {
|
||||||
|
if r == nil || r.db == nil {
|
||||||
|
return errors.New("repository not initialized")
|
||||||
|
}
|
||||||
|
if len(inputs) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *Repository) DeleteNftCounterStatesByForward(forwardID int64) error {
|
||||||
|
if r == nil || r.db == nil {
|
||||||
|
return errors.New("repository not initialized")
|
||||||
|
}
|
||||||
|
return r.db.Where("forward_id = ?", forwardID).Delete(&model.NftCounterState{}).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
func nftCounterStateFromInput(input NftCounterStateInput, now int64) (model.NftCounterState, bool) {
|
||||||
|
protocol := strings.ToLower(strings.TrimSpace(input.Protocol))
|
||||||
|
direction := strings.ToLower(strings.TrimSpace(input.Direction))
|
||||||
|
if input.NodeID <= 0 || input.ForwardID <= 0 || !isValidNftCounterProtocol(protocol) || !isValidNftCounterDirection(direction) {
|
||||||
|
return model.NftCounterState{}, false
|
||||||
|
}
|
||||||
|
if input.Bytes > uint64(math.MaxInt64) || input.Packets > uint64(math.MaxInt64) {
|
||||||
|
return model.NftCounterState{}, false
|
||||||
|
}
|
||||||
|
return model.NftCounterState{
|
||||||
|
NodeID: input.NodeID,
|
||||||
|
ForwardID: input.ForwardID,
|
||||||
|
Protocol: protocol,
|
||||||
|
Direction: direction,
|
||||||
|
RuleHash: strings.TrimSpace(input.RuleHash),
|
||||||
|
Bytes: int64(input.Bytes),
|
||||||
|
Packets: int64(input.Packets),
|
||||||
|
CollectedTime: input.CollectedTime,
|
||||||
|
CreatedTime: now,
|
||||||
|
UpdatedTime: now,
|
||||||
|
}, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func isValidNftCounterProtocol(protocol string) bool {
|
||||||
|
return protocol == nftCounterProtocolTCP || protocol == nftCounterProtocolUDP
|
||||||
|
}
|
||||||
|
|
||||||
|
func isValidNftCounterDirection(direction string) bool {
|
||||||
|
return direction == nftCounterDirectionToTarget || direction == nftCounterDirectionFromTarget
|
||||||
|
}
|
||||||
@@ -0,0 +1,161 @@
|
|||||||
|
package repo
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNftCounterStateUpsertUpdatesExistingKey(t *testing.T) {
|
||||||
|
r, err := Open(":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open repo: %v", err)
|
||||||
|
}
|
||||||
|
defer r.Close()
|
||||||
|
|
||||||
|
first := []NftCounterStateInput{
|
||||||
|
{
|
||||||
|
NodeID: 11,
|
||||||
|
ForwardID: 42,
|
||||||
|
Protocol: "tcp",
|
||||||
|
Direction: "to-target",
|
||||||
|
RuleHash: "hash-a",
|
||||||
|
Bytes: 100,
|
||||||
|
Packets: 10,
|
||||||
|
CollectedTime: 1000,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
NodeID: 0,
|
||||||
|
ForwardID: 42,
|
||||||
|
Protocol: "tcp",
|
||||||
|
Direction: "to-target",
|
||||||
|
Bytes: 999,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := r.UpsertNftCounterStates(first, 2000); err != nil {
|
||||||
|
t.Fatalf("first UpsertNftCounterStates: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
second := []NftCounterStateInput{
|
||||||
|
{
|
||||||
|
NodeID: 11,
|
||||||
|
ForwardID: 42,
|
||||||
|
Protocol: "tcp",
|
||||||
|
Direction: "to-target",
|
||||||
|
RuleHash: "hash-b",
|
||||||
|
Bytes: 250,
|
||||||
|
Packets: 25,
|
||||||
|
CollectedTime: 3000,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := r.UpsertNftCounterStates(second, 4000); err != nil {
|
||||||
|
t.Fatalf("second UpsertNftCounterStates: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := r.GetNftCounterStatesByNode(11)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||||
|
}
|
||||||
|
if len(rows) != 1 {
|
||||||
|
t.Fatalf("expected one counter state row, got %d: %+v", len(rows), rows)
|
||||||
|
}
|
||||||
|
got := rows[0]
|
||||||
|
if got.ForwardID != 42 || got.Protocol != "tcp" || got.Direction != "to-target" {
|
||||||
|
t.Fatalf("unexpected counter state key: %+v", got)
|
||||||
|
}
|
||||||
|
if got.RuleHash != "hash-b" || got.Bytes != 250 || got.Packets != 25 || got.CollectedTime != 3000 {
|
||||||
|
t.Fatalf("counter state was not updated: %+v", got)
|
||||||
|
}
|
||||||
|
if got.CreatedTime != 2000 || got.UpdatedTime != 4000 {
|
||||||
|
t.Fatalf("unexpected timestamps after upsert: %+v", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNftCounterStateDeleteByForwardRemovesOnlyMatchingRows(t *testing.T) {
|
||||||
|
r, err := Open(":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open repo: %v", err)
|
||||||
|
}
|
||||||
|
defer r.Close()
|
||||||
|
|
||||||
|
inputs := []NftCounterStateInput{
|
||||||
|
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "a", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||||
|
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "b", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||||
|
{NodeID: 12, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "c", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||||
|
}
|
||||||
|
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||||
|
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||||
|
}
|
||||||
|
if err := r.DeleteNftCounterStatesByForward(42); err != nil {
|
||||||
|
t.Fatalf("DeleteNftCounterStatesByForward: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
node11, err := r.GetNftCounterStatesByNode(11)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetNftCounterStatesByNode(11): %v", err)
|
||||||
|
}
|
||||||
|
if len(node11) != 1 || node11[0].ForwardID != 43 {
|
||||||
|
t.Fatalf("expected only forward 43 for node 11, got %+v", node11)
|
||||||
|
}
|
||||||
|
node12, err := r.GetNftCounterStatesByNode(12)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetNftCounterStatesByNode(12): %v", err)
|
||||||
|
}
|
||||||
|
if len(node12) != 0 {
|
||||||
|
t.Fatalf("expected forward 42 state removed from node 12, got %+v", node12)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNftCounterStateUpsertSkipsInvalidProtocolAndDirection(t *testing.T) {
|
||||||
|
r, err := Open(":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open repo: %v", err)
|
||||||
|
}
|
||||||
|
defer r.Close()
|
||||||
|
|
||||||
|
inputs := []NftCounterStateInput{
|
||||||
|
{NodeID: 11, ForwardID: 42, Protocol: "icmp", Direction: "to-target", RuleHash: "bad-protocol", Bytes: 100, Packets: 10, CollectedTime: 1000},
|
||||||
|
{NodeID: 11, ForwardID: 43, Protocol: "tcp", Direction: "sideways", RuleHash: "bad-direction", Bytes: 200, Packets: 20, CollectedTime: 1000},
|
||||||
|
{NodeID: 11, ForwardID: 44, Protocol: " UDP ", Direction: " FROM-TARGET ", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||||
|
}
|
||||||
|
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||||
|
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := r.GetNftCounterStatesByNode(11)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||||
|
}
|
||||||
|
if len(rows) != 1 {
|
||||||
|
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||||
|
}
|
||||||
|
if rows[0].ForwardID != 44 || rows[0].Protocol != "udp" || rows[0].Direction != "from-target" {
|
||||||
|
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNftCounterStateUpsertSkipsCountersAboveInt64(t *testing.T) {
|
||||||
|
r, err := Open(":memory:")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open repo: %v", err)
|
||||||
|
}
|
||||||
|
defer r.Close()
|
||||||
|
|
||||||
|
inputs := []NftCounterStateInput{
|
||||||
|
{NodeID: 11, ForwardID: 42, Protocol: "tcp", Direction: "to-target", RuleHash: "too-large", Bytes: uint64(math.MaxInt64) + 1, Packets: 10, CollectedTime: 1000},
|
||||||
|
{NodeID: 11, ForwardID: 43, Protocol: "udp", Direction: "from-target", RuleHash: "valid", Bytes: 300, Packets: 30, CollectedTime: 1000},
|
||||||
|
}
|
||||||
|
if err := r.UpsertNftCounterStates(inputs, 2000); err != nil {
|
||||||
|
t.Fatalf("UpsertNftCounterStates: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
rows, err := r.GetNftCounterStatesByNode(11)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetNftCounterStatesByNode: %v", err)
|
||||||
|
}
|
||||||
|
if len(rows) != 1 {
|
||||||
|
t.Fatalf("expected only the valid counter state row, got %d: %+v", len(rows), rows)
|
||||||
|
}
|
||||||
|
if rows[0].ForwardID != 43 || rows[0].Bytes != 300 || rows[0].Packets != 30 {
|
||||||
|
t.Fatalf("unexpected valid counter state row: %+v", rows[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user