Match any listed port or protocol instead of AND-ing them
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:41:42 +10:00
parent 410109515e
commit 9976bc9190
2 changed files with 286 additions and 164 deletions
+192 -164
View File
@@ -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)
+94
View File
@@ -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)
}
})
}
}