Pick reject type by resolved protocol number
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful

This commit is contained in:
2026-10-03 20:51:25 +10:00
parent 852d6bf2ca
commit 457056d58a
2 changed files with 40 additions and 27 deletions
+9 -7
View File
@@ -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}}
}
+31 -20
View File
@@ -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])
}
}
})
}
}