diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 161f944..2fc1e4d 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -428,6 +428,11 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, return err } + ip := net.ParseIP(dnatAddr) + if ip == nil { + return fmt.Errorf("invalid DNAT address %q", dnatAddr) + } + for _, srcIface := range srcIfaces { for _, m := range matches { var exprs []expr.Any @@ -450,11 +455,6 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, exprs = append(exprs, buildLog(logLevel, tag)...) } - ip := net.ParseIP(dnatAddr) - if ip == nil { - return fmt.Errorf("invalid DNAT address %q", dnatAddr) - } - if action == config.RuleRedirect { if dnatPort > 0 { portBytes := make([]byte, 2) @@ -975,8 +975,19 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") parseD := parsePortOrRange if isICMP { + if len(protos) > 1 && len(dports) > 0 { + return nil, fmt.Errorf("dport %v is ambiguous with %s in proto list %q", dports, p, proto) + } parseD = matchICMPType } + var pm []expr.Any + if p != "" { + m, err := matchProto(p) + if err != nil { + return nil, err + } + pm = m + } dalts, err := portAlternatives(dports, parseD) if err != nil { return nil, err @@ -987,11 +998,7 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) } for _, d := range dalts { for _, sp := range salts { - var e []expr.Any - if p != "" { - e = append(e, matchProto(p)...) - } - e = append(append(e, d...), sp...) + e := append(append(append([]expr.Any{}, pm...), d...), sp...) out = append(out, l4Match{proto: p, exprs: e}) } } @@ -1044,7 +1051,7 @@ func matchIfaceName(input bool, name string) []expr.Any { } } -func matchProto(proto string) []expr.Any { +func matchProto(proto string) ([]expr.Any, error) { var protoNum byte switch strings.ToLower(proto) { case "tcp": @@ -1064,13 +1071,16 @@ func matchProto(proto string) []expr.Any { case "sctp": protoNum = unix.IPPROTO_SCTP default: - n, _ := strconv.Atoi(proto) + n, err := strconv.ParseUint(proto, 10, 8) + if err != nil { + return nil, fmt.Errorf("unknown protocol %q", proto) + } protoNum = byte(n) } return []expr.Any{ &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{protoNum}}, - } + }, nil } func matchDPort(port uint16) []expr.Any { diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 50b1263..939c4a9 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1719,6 +1719,12 @@ func TestCompile_ListErrors(t *testing.T) { {"leading empty proto", config.Rule{Proto: ",udp", DPort: config.PortSpec{"80"}}}, {"empty port element", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,"}}}, {"unknown icmp type", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus"}}}, + {"unknown icmp type with code", config.Rule{Proto: "icmp", DPort: config.PortSpec{"bogus/0"}}}, + {"invalid icmp code", config.Rule{Proto: "icmp", DPort: config.PortSpec{"destination-unreachable/x"}}}, + {"icmp code out of range", config.Rule{Proto: "icmp", DPort: config.PortSpec{"3/256"}}}, + {"unknown proto in list", config.Rule{Proto: "tcp,udpp", DPort: config.PortSpec{"53"}}}, + {"proto number out of range", config.Rule{Proto: "256"}}, + {"dport with icmp in proto list", config.Rule{Proto: "icmp,tcp", DPort: config.PortSpec{"80"}}}, {"ratelimit with port list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,443"}, RateLimit: "10/sec"}}, {"connlimit with proto list", config.Rule{Proto: "tcp,udp", DPort: config.PortSpec{"53"}, ConnLimit: "10"}}, } @@ -1734,18 +1740,72 @@ func TestCompile_ListErrors(t *testing.T) { } } -func TestCompile_ListExpansionDiffStable(t *testing.T) { - cfg := listCfg(func(c *config.Config) { - c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}} - }) - state, err := NewCompiler(cfg).Compile() +func TestCompile_ICMPList(t *testing.T) { + state, err := NewCompiler(listCfg(func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "icmp", DPort: config.PortSpec{"echo-request,echo-reply", "destination-unreachable/4"}}} + })).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) } - if n := len(taggedRules(state, "input", "rule:0")); n != 4 { - t.Fatalf("got %d rule:0 rules, want 4", n) + var got [][]byte + for _, r := range taggedRules(state, "input", "rule:0") { + var tc []byte + for i, e := range r.Exprs { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Len == 1 { + tc = append(tc, r.Exprs[i+1].(*expr.Cmp).Data[0]) + } + } + got = append(got, tc) } - if cs := computeDiff(state, state); !cs.Empty() { - t.Errorf("diff not empty:\n%s", cs.Summary()) + want := [][]byte{{8}, {0}, {3, 4}} + if !reflect.DeepEqual(got, want) { + t.Errorf("icmp type/code per rule = %v, want %v", got, want) + } +} + +func TestCompile_ColonRanges(t *testing.T) { + tests := []struct { + name string + mod func(*config.Config) + chain string + tag string + offset uint32 + }{ + {"rule sport", func(c *config.Config) { + c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", SPort: config.PortSpec{"1024:2048"}}} + }, "input", "rule:0", 0}, + {"snat sport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "udp", SPort: config.PortSpec{"1024:2048"}}} + }, "postrouting", "snat:0", 0}, + {"snat dport", func(c *config.Config) { + c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp", DPort: config.PortSpec{"1024:2048"}}} + }, "postrouting", "snat:0", 2}, + {"conntrack dport", func(c *config.Config) { + c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"1024:2048"}}} + }, "prerouting", "conntrack:0:prerouting", 2}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + state, err := NewCompiler(listCfg(tt.mod)).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + rules := taggedRules(state, tt.chain, tt.tag) + if len(rules) != 1 { + t.Fatalf("got %d %s rules, want 1", len(rules), tt.tag) + } + var got []string + ex := rules[0].Exprs + for i, e := range ex { + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Offset == tt.offset && p.Len == 2 && i+2 < len(ex) { + lo, hi := ex[i+1].(*expr.Cmp), ex[i+2].(*expr.Cmp) + got = append(got, fmt.Sprintf("%d>=%d,%d<=%d", lo.Op, binary.BigEndian.Uint16(lo.Data), hi.Op, binary.BigEndian.Uint16(hi.Data))) + } + } + want := []string{fmt.Sprintf("%d>=1024,%d<=2048", expr.CmpOpGte, expr.CmpOpLte)} + if !reflect.DeepEqual(got, want) { + t.Errorf("range match = %v, want %v", got, want) + } + }) } }