Pick reject type by resolved protocol number
This commit is contained in:
@@ -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}}
|
||||
}
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user