From 457056d58a7531fac1e4e93f8e533fde7311bba4 Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sat, 3 Oct 2026 20:51:25 +1000 Subject: [PATCH] Pick reject type by resolved protocol number --- internal/nftables/compiler.go | 16 ++++++---- internal/nftables/compiler_test.go | 51 ++++++++++++++++++------------ 2 files changed, 40 insertions(+), 27 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 9bd47eb..2994a80 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -955,7 +955,7 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor } type l4Match struct { - proto string + proto byte exprs []expr.Any } @@ -973,9 +973,11 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) return nil, fmt.Errorf("empty element in proto list %q", proto) } var pm []expr.Any + var n byte isICMP := false if p != "" { - n, err := protoNumber(p) + var err error + n, err = protoNumber(p) if err != nil { return nil, err } @@ -1003,7 +1005,7 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) for _, d := range dalts { for _, sp := range salts { e := append(append(append([]expr.Any{}, pm...), d...), sp...) - out = append(out, l4Match{proto: p, exprs: e}) + out = append(out, l4Match{proto: n, exprs: e}) } } } @@ -1665,7 +1667,7 @@ func logLevelToNF(level string) expr.LogLevel { } } -func actionVerdict(action config.RuleAction, proto string, family config.AddressFamily) []expr.Any { +func actionVerdict(action config.RuleAction, proto byte, family config.AddressFamily) []expr.Any { switch action { case config.RuleAccept: return []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}} @@ -1690,8 +1692,8 @@ func actionVerdict(action config.RuleAction, proto string, family config.Address } } -func rejectExprs(proto string, family config.AddressFamily) []expr.Any { - if strings.ToLower(proto) == "tcp" { +func rejectExprs(proto byte, family config.AddressFamily) []expr.Any { + if proto == unix.IPPROTO_TCP { return []expr.Any{&expr.Reject{ Type: unix.NFT_REJECT_TCP_RST, Code: 0, @@ -1710,7 +1712,7 @@ func policyVerdict(action config.PolicyAction, family config.AddressFamily) []ex case config.PolicyDrop: return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} case config.PolicyReject: - return rejectExprs("", family) + return rejectExprs(0, family) default: return []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 912ce65..1a221b5 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -978,7 +978,7 @@ func TestNegatedAddress(t *testing.T) { } func TestRejectTCPRST(t *testing.T) { - exprs := rejectExprs("tcp", config.FamilyINET) + exprs := rejectExprs(unix.IPPROTO_TCP, config.FamilyINET) if len(exprs) != 1 { t.Fatalf("expected 1 expr, got %d", len(exprs)) } @@ -987,7 +987,7 @@ func TestRejectTCPRST(t *testing.T) { t.Errorf("TCP reject should use NFT_REJECT_TCP_RST (1), got %d", rej.Type) } - exprs = rejectExprs("udp", config.FamilyINET) + exprs = rejectExprs(unix.IPPROTO_UDP, config.FamilyINET) rej = exprs[0].(*expr.Reject) if rej.Type != 2 { t.Errorf("non-TCP reject should use NFT_REJECT_ICMPX_UNREACH (2), got %d", rej.Type) @@ -1705,25 +1705,36 @@ func TestCompile_ListExpansionCounts(t *testing.T) { } func TestCompile_RejectPerProto(t *testing.T) { - state, err := NewCompiler(listCfg(func(c *config.Config) { - c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}} - })).Compile() - if err != nil { - t.Fatalf("Compile() error: %v", err) + tests := []struct { + proto string + want []uint32 + }{ + {"tcp,udp", []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}}, + {"6", []uint32{unix.NFT_REJECT_TCP_RST}}, + {"6,17", []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}}, } - rules := taggedRules(state, "input", "rule:0") - if len(rules) != 2 { - t.Fatalf("got %d rules, want 2", len(rules)) - } - want := []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH} - for i, r := range rules { - rej, ok := r.Exprs[len(r.Exprs)-1].(*expr.Reject) - if !ok { - t.Fatalf("rule %d: last expr %T, want *expr.Reject", i, r.Exprs[len(r.Exprs)-1]) - } - if rej.Type != want[i] { - t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, want[i]) - } + for _, tt := range tests { + t.Run(tt.proto, func(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: tt.proto, DPort: config.PortSpec{"53"}}} + })).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + rules := taggedRules(state, "input", "rule:0") + if len(rules) != len(tt.want) { + t.Fatalf("got %d rules, want %d", len(rules), len(tt.want)) + } + for i, r := range rules { + rej, ok := r.Exprs[len(r.Exprs)-1].(*expr.Reject) + if !ok { + t.Fatalf("rule %d: last expr %T, want *expr.Reject", i, r.Exprs[len(r.Exprs)-1]) + } + if rej.Type != tt.want[i] { + t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, tt.want[i]) + } + } + }) } }