diff --git a/go-backend/internal/store/model/model.go b/go-backend/internal/store/model/model.go index e655950..6ac19ad 100644 --- a/go-backend/internal/store/model/model.go +++ b/go-backend/internal/store/model/model.go @@ -131,6 +131,22 @@ type NftRuleBinding struct { 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 { ID int64 `gorm:"primaryKey;autoIncrement"` Name string `gorm:"type:varchar(100);not null"` diff --git a/go-backend/internal/store/repo/repository.go b/go-backend/internal/store/repo/repository.go index 2b74377..46c9cc7 100644 --- a/go-backend/internal/store/repo/repository.go +++ b/go-backend/internal/store/repo/repository.go @@ -282,6 +282,7 @@ func autoMigrateAll(db *gorm.DB) error { &model.Node{}, &model.NodeSSHConfig{}, &model.NftRuleBinding{}, + &model.NftCounterState{}, &model.SpeedLimit{}, &model.StatisticsFlow{}, &model.Tunnel{}, diff --git a/go-backend/internal/store/repo/repository_nft_counter.go b/go-backend/internal/store/repo/repository_nft_counter.go new file mode 100644 index 0000000..66b30d4 --- /dev/null +++ b/go-backend/internal/store/repo/repository_nft_counter.go @@ -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 +} diff --git a/go-backend/internal/store/repo/repository_nft_counter_test.go b/go-backend/internal/store/repo/repository_nft_counter_test.go new file mode 100644 index 0000000..d955843 --- /dev/null +++ b/go-backend/internal/store/repo/repository_nft_counter_test.go @@ -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]) + } +}