diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 45e6549..2af98ad 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -220,34 +220,30 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { chains = []string{"prerouting", "output"} } + matches, err := l4Matches(ct.Proto, ct.DPort, nil) + if err != nil { + return fmt.Errorf("conntrack[%d]: %w", i, err) + } + for _, chain := range chains { - var exprs []expr.Any + for _, m := range matches { + exprs := append([]expr.Any{}, m.exprs...) - if ct.Proto != "" { - exprs = append(exprs, matchProto(ct.Proto)...) - } - for _, p := range ct.DPort { - pe, err := parsePortOrRange(p) - if err != nil { - return fmt.Errorf("conntrack[%d]: %w", i, err) + switch ct.Action { + case config.ConntrackNoTrack: + exprs = append(exprs, &expr.Notrack{}) + case config.ConntrackHelper: + continue + case config.ConntrackDrop: + exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) } - exprs = append(exprs, pe...) - } - switch ct.Action { - case config.ConntrackNoTrack: - exprs = append(exprs, &expr.Notrack{}) - case config.ConntrackHelper: - continue - case config.ConntrackDrop: - exprs = append(exprs, &expr.Verdict{Kind: expr.VerdictDrop}) + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag + ":" + chain, + }) } - - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag + ":" + chain, - }) } } return nil @@ -362,29 +358,33 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p for _, srcIface := range srcIfaces { for _, dstIface := range dstIfaces { - exprs, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr) + matches, err := c.buildMatchExprs(srcIface, dstIface, chain, proto, dports, sports, srcAddr, dstAddr) if err != nil { return err } - if section != "" && section != config.SectionAll { - exprs = append(exprs, matchSection(section)...) - } + for _, m := range matches { + exprs := m.exprs - if logLevel != "" { - exprs = append(exprs, buildLog(logLevel, tag)...) - } + if section != "" && section != config.SectionAll { + exprs = append(exprs, matchSection(section)...) + } - verdict := actionVerdict(action, proto, c.cfg.Settings.AddressFamily) - if verdict != nil { - exprs = append(exprs, verdict...) - } + if logLevel != "" { + exprs = append(exprs, buildLog(logLevel, tag)...) + } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag, - }) + verdict := actionVerdict(action, m.proto, c.cfg.Settings.AddressFamily) + if verdict != nil { + exprs = append(exprs, verdict...) + } + + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag, + }) + } } } return nil @@ -413,102 +413,99 @@ func (c *Compiler) compileDNATRule(state *FirewallState, tag, srcSpec, dstSpec, srcIfaces := c.resolveZoneInterfaces(srcZone) + matches, err := l4Matches(proto, dports, nil) + if err != nil { + return err + } + for _, srcIface := range srcIfaces { - var exprs []expr.Any + for _, m := range matches { + var exprs []expr.Any - if srcIface != "" { - exprs = append(exprs, matchIfaceName(true, srcIface)...) - } - - if srcAddr != "" { - src, err := matchSourceCIDR(srcAddr) - if err != nil { - return err + if srcIface != "" { + exprs = append(exprs, matchIfaceName(true, srcIface)...) } - exprs = append(exprs, src...) - } - if proto != "" { - exprs = append(exprs, matchProto(proto)...) - } - - for _, portStr := range dports { - pe, err := parsePortOrRange(portStr) - if err != nil { - return err - } - exprs = append(exprs, pe...) - } - - if logLevel != "" { - 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) - binary.BigEndian.PutUint16(portBytes, dnatPort) - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: portBytes}, - &expr.Redir{RegisterProtoMin: 1}, - ) - } else { - exprs = append(exprs, &expr.Redir{}) - } - } else { - if ip4 := ip.To4(); ip4 != nil { - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: ip4}, - ) - natExpr := &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: unix.NFPROTO_IPV4, - RegAddrMin: 1, - RegAddrMax: 1, + if srcAddr != "" { + src, err := matchSourceCIDR(srcAddr) + if err != nil { + return err } + exprs = append(exprs, src...) + } + + exprs = append(exprs, m.exprs...) + + if logLevel != "" { + 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) binary.BigEndian.PutUint16(portBytes, dnatPort) exprs = append(exprs, - &expr.Immediate{Register: 2, Data: portBytes}, + &expr.Immediate{Register: 1, Data: portBytes}, + &expr.Redir{RegisterProtoMin: 1}, ) - natExpr.RegProtoMin = 2 - natExpr.RegProtoMax = 2 + } else { + exprs = append(exprs, &expr.Redir{}) } - exprs = append(exprs, natExpr) } else { - exprs = append(exprs, - &expr.Immediate{Register: 1, Data: ip.To16()}, - ) - natExpr := &expr.NAT{ - Type: expr.NATTypeDestNAT, - Family: unix.NFPROTO_IPV6, - RegAddrMin: 1, - RegAddrMax: 1, - } - if dnatPort > 0 { - portBytes := make([]byte, 2) - binary.BigEndian.PutUint16(portBytes, dnatPort) + if ip4 := ip.To4(); ip4 != nil { exprs = append(exprs, - &expr.Immediate{Register: 2, Data: portBytes}, + &expr.Immediate{Register: 1, Data: ip4}, ) - natExpr.RegProtoMin = 2 - natExpr.RegProtoMax = 2 + natExpr := &expr.NAT{ + Type: expr.NATTypeDestNAT, + Family: unix.NFPROTO_IPV4, + RegAddrMin: 1, + RegAddrMax: 1, + } + if dnatPort > 0 { + portBytes := make([]byte, 2) + binary.BigEndian.PutUint16(portBytes, dnatPort) + exprs = append(exprs, + &expr.Immediate{Register: 2, Data: portBytes}, + ) + natExpr.RegProtoMin = 2 + natExpr.RegProtoMax = 2 + } + exprs = append(exprs, natExpr) + } else { + exprs = append(exprs, + &expr.Immediate{Register: 1, Data: ip.To16()}, + ) + natExpr := &expr.NAT{ + Type: expr.NATTypeDestNAT, + Family: unix.NFPROTO_IPV6, + RegAddrMin: 1, + RegAddrMax: 1, + } + if dnatPort > 0 { + portBytes := make([]byte, 2) + binary.BigEndian.PutUint16(portBytes, dnatPort) + exprs = append(exprs, + &expr.Immediate{Register: 2, Data: portBytes}, + ) + natExpr.RegProtoMin = 2 + natExpr.RegProtoMax = 2 + } + exprs = append(exprs, natExpr) } - exprs = append(exprs, natExpr) } - } - state.Rules[chain] = append(state.Rules[chain], ManagedRule{ - Chain: chain, - Exprs: exprs, - Tag: tag, - }) + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: exprs, + Tag: tag, + }) + } } return nil } @@ -588,25 +585,12 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { exprs = append(exprs, srcExprs...) } - if snat.Proto != "" { - exprs = append(exprs, matchProto(snat.Proto)...) - } - - for _, portStr := range snat.DPort { - pe, err := parsePortOrRange(portStr) - if err != nil { - return fmt.Errorf("snat[%d] dport: %w", i, err) - } - exprs = append(exprs, pe...) - } - - for _, portStr := range snat.SPort { - pe, err := parseSPortOrRange(portStr) - if err != nil { - return fmt.Errorf("snat[%d] sport: %w", i, err) - } - exprs = append(exprs, pe...) + matches, err := l4Matches(snat.Proto, snat.DPort, snat.SPort) + if err != nil { + return fmt.Errorf("snat[%d]: %w", i, err) } + head := exprs + exprs = nil if snat.Mark != "" { exprs = append(exprs, matchMark(snat.Mark)...) @@ -652,11 +636,13 @@ func (c *Compiler) compileSNAT(state *FirewallState) error { } } - state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ - Chain: "postrouting", - Exprs: exprs, - Tag: tag, - }) + for _, m := range matches { + state.Rules["postrouting"] = append(state.Rules["postrouting"], ManagedRule{ + Chain: "postrouting", + Exprs: append(append(append([]expr.Any{}, head...), m.exprs...), exprs...), + Tag: tag, + }) + } } return nil @@ -922,7 +908,7 @@ func (c *Compiler) expandZoneRef(ref string) []string { return []string{base} } -func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]expr.Any, error) { +func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dports, sports config.PortSpec, srcAddr, dstAddr string) ([]l4Match, error) { var exprs []expr.Any if srcIface != "" { @@ -948,34 +934,76 @@ func (c *Compiler) buildMatchExprs(srcIface, dstIface, chain, proto string, dpor exprs = append(exprs, dst...) } + matches, err := l4Matches(proto, dports, sports) + if err != nil { + return nil, err + } + for i := range matches { + matches[i].exprs = append(append([]expr.Any{}, exprs...), matches[i].exprs...) + } + return matches, nil +} + +type l4Match struct { + proto string + exprs []expr.Any +} + +// l4Matches yields one alternative per (proto, dport, sport): nft ANDs a rule's exprs, so lists need one rule each. +func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error) { + protos := []string{""} if proto != "" { - exprs = append(exprs, matchProto(proto)...) + protos = strings.Split(proto, ",") } - isICMP := strings.EqualFold(proto, "icmp") || strings.EqualFold(proto, "icmpv6") || strings.EqualFold(proto, "ipv6-icmp") - - for _, portStr := range dports { + var out []l4Match + for _, p := range protos { + p = strings.TrimSpace(p) + isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp") + parseD := parsePortOrRange if isICMP { - pe := matchICMPType(portStr) - exprs = append(exprs, pe...) - } else { - pe, err := parsePortOrRange(portStr) - if err != nil { - return nil, err - } - exprs = append(exprs, pe...) + parseD = func(s string) ([]expr.Any, error) { return matchICMPType(s), nil } } - } - - for _, portStr := range sports { - pe, err := parseSPortOrRange(portStr) + dalts, err := portAlternatives(dports, parseD) if err != nil { return nil, err } - exprs = append(exprs, pe...) + salts, err := portAlternatives(sports, parseSPortOrRange) + if err != nil { + return nil, err + } + 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...) + out = append(out, l4Match{proto: p, exprs: e}) + } + } } + return out, nil +} - return exprs, nil +func portAlternatives(ports config.PortSpec, parse func(string) ([]expr.Any, error)) ([][]expr.Any, error) { + var alts [][]expr.Any + for _, item := range ports { + for _, s := range strings.Split(item, ",") { + if s = strings.TrimSpace(s); s == "" { + continue + } + e, err := parse(s) + if err != nil { + return nil, err + } + alts = append(alts, e) + } + } + if len(alts) == 0 { + return [][]expr.Any{nil}, nil + } + return alts, nil } // matchIfaceName matches an interface name, supporting wildcard "+" suffix. @@ -1218,8 +1246,8 @@ func matchConnLimit(spec string) []expr.Any { } func parsePortOrRange(s string) ([]expr.Any, error) { - if strings.Contains(s, "-") { - parts := strings.SplitN(s, "-", 2) + if i := strings.IndexAny(s, "-:"); i >= 0 { + parts := []string{s[:i], s[i+1:]} low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid port range low %q: %w", parts[0], err) @@ -1238,8 +1266,8 @@ func parsePortOrRange(s string) ([]expr.Any, error) { } func parseSPortOrRange(s string) ([]expr.Any, error) { - if strings.Contains(s, "-") { - parts := strings.SplitN(s, "-", 2) + if i := strings.IndexAny(s, "-:"); i >= 0 { + parts := []string{s[:i], s[i+1:]} low, err := strconv.ParseUint(parts[0], 10, 16) if err != nil { return nil, fmt.Errorf("invalid sport range low %q: %w", parts[0], err) diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 84fa751..bcc1db7 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -1,6 +1,10 @@ package nftables import ( + "encoding/binary" + "fmt" + "reflect" + "strings" "testing" "github.com/google/nftables/expr" @@ -1531,3 +1535,93 @@ func TestCompile_PolicyRateLimit(t *testing.T) { } t.Error("policy:0 not found in input chain") } + +func TestCompile_PortAndProtoLists(t *testing.T) { + type want struct { + proto byte + dport string // "80" for an exact compare, "8000-8100" for a range + } + tests := []struct { + name string + rule config.Rule + chain string + want []want + }{ + { + name: "multi-port", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80", "443"}}, + chain: "input", + want: []want{{6, "80"}, {6, "443"}}, + }, + { + name: "range in list", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22", "8000:8100"}}, + chain: "input", + want: []want{{6, "22"}, {6, "8000-8100"}}, + }, + { + name: "comma string", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80,443"}}, + chain: "input", + want: []want{{6, "80"}, {6, "443"}}, + }, + { + name: "tcp,udp", + rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}, + chain: "input", + want: []want{{6, "53"}, {17, "53"}}, + }, + { + name: "dnat tcp,udp", + rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.10", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}, + chain: "prerouting", + want: []want{{6, "53"}, {17, "53"}}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, + Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}}, + Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}}, + Policy: []config.Policy{{Source: "all", Dest: "all", Action: config.PolicyDrop}}, + Rules: []config.Rule{tt.rule}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + + var got []want + for _, r := range state.Rules[tt.chain] { + if r.Tag != "rule:0" { + continue + } + var w want + var portCmps []string + for i, e := range r.Exprs { + if m, ok := e.(*expr.Meta); ok && m.Key == expr.MetaKeyL4PROTO { + w.proto = r.Exprs[i+1].(*expr.Cmp).Data[0] + } + if p, ok := e.(*expr.Payload); ok && p.Base == expr.PayloadBaseTransportHeader && p.Offset == 2 { + for _, c := range r.Exprs[i+1:] { + cmp, ok := c.(*expr.Cmp) + if !ok { + break + } + portCmps = append(portCmps, fmt.Sprint(binary.BigEndian.Uint16(cmp.Data))) + } + } + } + w.dport = strings.Join(portCmps, "-") + got = append(got, w) + } + + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("rules = %+v, want %+v", got, tt.want) + } + }) + } +}