Match any listed port or protocol instead of AND-ing them
This commit is contained in:
+192
-164
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user