Files
tomswall/internal/nftables/compiler_test.go
T
unkin-agent 445d14c61e
ci/woodpecker/pr/pre-commit Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/test Pipeline was successful
bump google/nftables to v0.3.0
Set NAT.Specified on DNAT with a port so compiled exprs match kernel readback.
2026-10-03 22:56:20 +10:00

2420 lines
66 KiB
Go

package nftables
import (
"bytes"
"encoding/binary"
"fmt"
"log/slog"
"net"
"reflect"
"strings"
"testing"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
"git.unkin.net/unkin/tomswall/internal/config"
)
func TestSplitZoneSpec(t *testing.T) {
tests := []struct {
input string
wantZone string
wantAddr string
}{
{"net", "net", ""},
{"net:192.168.1.0/24", "net", "192.168.1.0/24"},
{"loc:10.0.0.1", "loc", "10.0.0.1"},
{"fw", "fw", ""},
{"all", "all", ""},
{"dmz:2001:db8::/32", "dmz", "2001:db8::/32"},
{"", "", ""},
}
for _, tt := range tests {
zone, addr := splitZoneSpec(tt.input)
if zone != tt.wantZone || addr != tt.wantAddr {
t.Errorf("splitZoneSpec(%q) = (%q, %q), want (%q, %q)",
tt.input, zone, addr, tt.wantZone, tt.wantAddr)
}
}
}
func TestParsePort(t *testing.T) {
tests := []struct {
input string
want uint16
wantErr bool
}{
{"22", 22, false},
{"80", 80, false},
{"443", 443, false},
{"65535", 65535, false},
{"0", 0, false},
{"1", 1, false},
{"65536", 0, true},
{"-1", 0, true},
{"abc", 0, true},
{"", 0, true},
{"99999", 0, true},
}
for _, tt := range tests {
got, err := parsePort(tt.input)
if tt.wantErr {
if err == nil {
t.Errorf("parsePort(%q) = %d, want error", tt.input, got)
}
continue
}
if err != nil {
t.Errorf("parsePort(%q) returned error: %v", tt.input, err)
continue
}
if got != tt.want {
t.Errorf("parsePort(%q) = %d, want %d", tt.input, got, tt.want)
}
}
}
func TestNewCompiler(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{
TableName: "test",
LogLevel: "info",
},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
if c == nil {
t.Fatal("NewCompiler returned nil")
}
if c.cfg != cfg {
t.Error("compiler cfg does not match input cfg")
}
}
func TestCompiler_SelectChain(t *testing.T) {
cfg := &config.Config{
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
"loc": {Type: config.ZoneIP},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
tests := []struct {
src, dst, fw string
want string
}{
{"net", "fw", "fw", "input"},
{"fw", "net", "fw", "output"},
{"net", "loc", "fw", "forward"},
{"loc", "net", "fw", "forward"},
}
for _, tt := range tests {
got := c.selectChain(tt.src, tt.dst, tt.fw)
if got != tt.want {
t.Errorf("selectChain(%q, %q, %q) = %q, want %q",
tt.src, tt.dst, tt.fw, got, tt.want)
}
}
}
func TestCompiler_ResolveZoneInterfaces(t *testing.T) {
cfg := &config.Config{
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
"loc": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "loc", Interface: "eth1"},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
ifaces := c.resolveZoneInterfaces("net", "")
if len(ifaces) != 1 || ifaces[0] != "eth0" {
t.Errorf("resolveZoneInterfaces(net) = %v, want [eth0]", ifaces)
}
ifaces = c.resolveZoneInterfaces("all", "")
if len(ifaces) != 1 || ifaces[0] != "" {
t.Errorf("resolveZoneInterfaces(all) = %v, want [\"\"]", ifaces)
}
ifaces = c.resolveZoneInterfaces("fw", "")
if len(ifaces) != 1 || ifaces[0] != "" {
t.Errorf("resolveZoneInterfaces(fw) = %v, want [\"\"]", ifaces)
}
}
func TestCompiler_ExpandZoneRef(t *testing.T) {
cfg := &config.Config{
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
"loc": {Type: config.ZoneIP},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
zones := c.expandZoneRef("net")
if len(zones) != 1 || zones[0] != "net" {
t.Errorf("expandZoneRef(net) = %v, want [net]", zones)
}
zones = c.expandZoneRef("all")
if len(zones) != 3 {
t.Errorf("expandZoneRef(all) = %v, want 3 zones", zones)
}
seen := make(map[string]bool)
for _, z := range zones {
seen[z] = true
}
for _, name := range []string{"fw", "net", "loc"} {
if !seen[name] {
t.Errorf("expandZoneRef(all) missing zone %q", name)
}
}
zones = c.expandZoneRef("all+")
if len(zones) != 3 {
t.Errorf("expandZoneRef(all+) = %v, want 3 zones", zones)
}
}
func TestMatchSourceCIDR_IPv6(t *testing.T) {
tests := []struct {
input string
wantLen int
wantErr bool
}{
{"192.168.1.0/24", 3, false},
{"10.0.0.1", 2, false},
{"fd10:10:9::/64", 3, false},
{"2001:db8::1", 2, false},
{"fd74:212::/48", 3, false},
{"invalid", 0, true},
}
for _, tt := range tests {
exprs, err := matchSourceCIDR(tt.input)
if tt.wantErr {
if err == nil {
t.Errorf("matchSourceCIDR(%q) should fail", tt.input)
}
continue
}
if err != nil {
t.Errorf("matchSourceCIDR(%q) error: %v", tt.input, err)
continue
}
if len(exprs) != tt.wantLen {
t.Errorf("matchSourceCIDR(%q) returned %d expressions, want %d", tt.input, len(exprs), tt.wantLen)
}
}
}
func TestMatchDestCIDR_IPv6(t *testing.T) {
tests := []struct {
input string
wantLen int
wantErr bool
}{
{"192.168.1.0/24", 3, false},
{"10.0.0.1", 2, false},
{"fd10:10:9::/64", 3, false},
{"2001:db8::1", 2, false},
{"invalid", 0, true},
}
for _, tt := range tests {
exprs, err := matchDestCIDR(tt.input)
if tt.wantErr {
if err == nil {
t.Errorf("matchDestCIDR(%q) should fail", tt.input)
}
continue
}
if err != nil {
t.Errorf("matchDestCIDR(%q) error: %v", tt.input, err)
continue
}
if len(exprs) != tt.wantLen {
t.Errorf("matchDestCIDR(%q) returned %d expressions, want %d", tt.input, len(exprs), tt.wantLen)
}
}
}
func TestParsePortOrRange(t *testing.T) {
tests := []struct {
input string
wantLen int
wantErr bool
}{
{"80", 2, false},
{"443", 2, false},
{"1024-65535", 3, false},
{"80-90", 3, false},
{"abc", 0, true},
{"80-abc", 0, true},
}
for _, tt := range tests {
exprs, err := parsePortOrRange(tt.input)
if tt.wantErr {
if err == nil {
t.Errorf("parsePortOrRange(%q) should fail", tt.input)
}
continue
}
if err != nil {
t.Errorf("parsePortOrRange(%q) error: %v", tt.input, err)
continue
}
if len(exprs) != tt.wantLen {
t.Errorf("parsePortOrRange(%q) returned %d expressions, want %d", tt.input, len(exprs), tt.wantLen)
}
}
}
func TestParseSPortOrRange(t *testing.T) {
tests := []struct {
input string
wantLen int
wantErr bool
}{
{"22", 2, false},
{"1024-65535", 3, false},
{"bad", 0, true},
}
for _, tt := range tests {
exprs, err := parseSPortOrRange(tt.input)
if tt.wantErr {
if err == nil {
t.Errorf("parseSPortOrRange(%q) should fail", tt.input)
}
continue
}
if err != nil {
t.Errorf("parseSPortOrRange(%q) error: %v", tt.input, err)
continue
}
if len(exprs) != tt.wantLen {
t.Errorf("parseSPortOrRange(%q) returned %d expressions, want %d", tt.input, len(exprs), tt.wantLen)
}
}
}
func TestMatchIfaceName_Wildcard(t *testing.T) {
exact := matchIfaceName(true, "eth0")
if len(exact) != 2 {
t.Fatalf("matchIfaceName(true, eth0) returned %d exprs, want 2", len(exact))
}
wild := matchIfaceName(true, "tun+")
if len(wild) != 2 {
t.Fatalf("matchIfaceName(true, tun+) returned %d exprs, want 2", len(wild))
}
}
func TestMatchCtState(t *testing.T) {
exprs := matchCtState(ctStateEstablished | ctStateRelated)
if len(exprs) != 3 {
t.Errorf("matchCtState returned %d expressions, want 3", len(exprs))
}
}
func TestCompile_ConntrackFastPath(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},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, chain := range []string{"input", "forward", "output"} {
rules := state.Rules[chain]
foundFastpath := false
foundInvalid := false
for _, r := range rules {
if r.Tag == "ct:fastpath:"+chain {
foundFastpath = true
}
if r.Tag == "ct:invalid:"+chain {
foundInvalid = true
}
}
if !foundFastpath {
t.Errorf("chain %q missing ct:fastpath rule", chain)
}
if !foundInvalid {
t.Errorf("chain %q missing ct:invalid rule", chain)
}
}
}
func TestCompile_IntraZone(t *testing.T) {
routeback := true
cfg := &config.Config{
Settings: config.Settings{
TableName: "test",
AddressFamily: config.FamilyINET,
},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"loc": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "loc", Interface: "eth1", Options: config.InterfaceOptions{RouteBack: &routeback}},
},
Policy: []config.Policy{
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
found := false
for _, r := range state.Rules["forward"] {
if r.Tag == "intra:loc:eth1" {
found = true
break
}
}
if !found {
t.Error("no intra-zone rule found for loc/eth1 in forward chain")
}
}
func TestCompile_DNAT(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},
"loc": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "loc", Interface: "eth1"},
},
Policy: []config.Policy{
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
Rules: []config.Rule{
{
Action: config.RuleDNAT,
Source: "net",
Dest: "loc:192.168.1.5:22",
Proto: "tcp",
DPort: config.PortSpec{"2222"},
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
var nat *expr.NAT
for _, r := range state.Rules["prerouting"] {
if r.Tag == "rule:0" {
nat, _ = r.Exprs[len(r.Exprs)-1].(*expr.NAT)
break
}
}
if nat == nil {
t.Fatal("no DNAT rule found in prerouting chain")
}
// The kernel reports PROTO_SPECIFIED whenever a port register is set.
if nat.RegProtoMin != 2 || !nat.Specified {
t.Errorf("DNAT with port must set RegProtoMin and Specified to match kernel readback, got %+v", nat)
}
}
func TestCompile_StaticNAT(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},
},
StaticNAT: []config.StaticNAT{
{
External: "203.0.113.10",
Interface: "eth0",
Internal: "192.168.1.10",
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
dnatFound := false
snatFound := false
for _, r := range state.Rules["prerouting"] {
if r.Tag == "staticnat:dnat:0" {
dnatFound = true
}
}
for _, r := range state.Rules["postrouting"] {
if r.Tag == "staticnat:snat:0" {
snatFound = true
}
}
if !dnatFound {
t.Error("no static NAT DNAT rule found in prerouting")
}
if !snatFound {
t.Error("no static NAT SNAT rule found in postrouting")
}
}
func TestCompile_Logging(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: "net", Dest: "all", Action: config.PolicyDrop, Log: "info"},
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
found := false
for _, r := range state.Rules["input"] {
if r.Tag == "policy:0" {
found = true
if len(r.Exprs) < 4 {
t.Errorf("policy:0 has %d exprs, want >= 4 (iface + log + verdict)", len(r.Exprs))
}
break
}
}
if !found {
t.Error("policy:0 not found in input chain")
}
}
func TestCompile_ConntrackNoTrack(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},
},
Conntrack: []config.ConntrackRule{
{
Action: config.ConntrackNoTrack,
Source: "net",
Dest: "fw",
Proto: "udp",
DPort: config.PortSpec{"53"},
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
found := false
for _, r := range state.Rules["prerouting"] {
if r.Tag == "conntrack:0:prerouting" {
found = true
break
}
}
if !found {
t.Error("no notrack rule found in prerouting chain")
}
}
func TestLogLevelToNF(t *testing.T) {
tests := []struct {
input string
want expr.LogLevel
}{
{"emerg", expr.LogLevelEmerg},
{"alert", expr.LogLevelAlert},
{"crit", expr.LogLevelCrit},
{"err", expr.LogLevelErr},
{"error", expr.LogLevelErr},
{"warn", expr.LogLevelWarning},
{"warning", expr.LogLevelWarning},
{"notice", expr.LogLevelNotice},
{"info", expr.LogLevelInfo},
{"debug", expr.LogLevelDebug},
{"unknown", expr.LogLevelWarning},
}
for _, tt := range tests {
got := logLevelToNF(tt.input)
if got != tt.want {
t.Errorf("logLevelToNF(%q) = %d, want %d", tt.input, got, tt.want)
}
}
}
func TestDiffEngine_DetectsModifications(t *testing.T) {
current := &FirewallState{
Rules: map[string][]ManagedRule{
"input": {
{Chain: "input", Tag: "rule:0", Exprs: []expr.Any{
&expr.Verdict{Kind: expr.VerdictAccept},
}},
},
},
}
desired := &FirewallState{
Rules: map[string][]ManagedRule{
"input": {
{Chain: "input", Tag: "rule:0", Exprs: []expr.Any{
&expr.Verdict{Kind: expr.VerdictDrop},
}},
},
},
}
cs := computeDiff(current, desired)
if len(cs.Remove) != 1 {
t.Errorf("expected 1 removal, got %d", len(cs.Remove))
}
if len(cs.Add) != 1 {
t.Errorf("expected 1 addition, got %d", len(cs.Add))
}
}
func TestDiffEngine_NoChangeWhenIdentical(t *testing.T) {
state := &FirewallState{
Rules: map[string][]ManagedRule{
"input": {
{Chain: "input", Tag: "rule:0", Exprs: []expr.Any{
&expr.Verdict{Kind: expr.VerdictAccept},
}},
},
},
}
cs := computeDiff(state, state)
if !cs.Empty() {
t.Errorf("expected empty changeset, got %d adds and %d removes", len(cs.Add), len(cs.Remove))
}
}
func TestCompile_SPortMatching(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{
{
Action: config.RuleAccept,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"22"},
SPort: config.PortSpec{"1024-65535"},
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
found := false
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
found = true
break
}
}
if !found {
t.Error("rule with sport not found in input chain")
}
}
func TestCompile_PortRange(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{
{
Action: config.RuleAccept,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"1024-65535"},
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
found := false
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
found = true
break
}
}
if !found {
t.Error("rule with port range not found")
}
}
func TestCompile_LogAction(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{
{
Action: config.RuleLog,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"22"},
Log: "info",
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
found := false
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
found = true
break
}
}
if !found {
t.Error("log rule not found")
}
}
func TestMatchSection(t *testing.T) {
tests := []struct {
section config.RuleSection
wantLen int
}{
{config.SectionEstablished, 3},
{config.SectionRelated, 3},
{config.SectionInvalid, 3},
{config.SectionUntracked, 3},
{config.SectionNew, 3},
{config.SectionAll, 0},
{"", 0},
}
for _, tt := range tests {
exprs := matchSection(tt.section)
if len(exprs) != tt.wantLen {
t.Errorf("matchSection(%q) returned %d exprs, want %d", tt.section, len(exprs), tt.wantLen)
}
}
}
func TestCompile_RuleSection(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{
{
Action: config.RuleAccept,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"22"},
Section: config.SectionEstablished,
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
if len(r.Exprs) < 6 {
t.Errorf("rule with section should have >= 6 exprs (iface+proto+dport+ctstate+verdict), got %d", len(r.Exprs))
}
return
}
}
t.Error("rule:0 not found in input chain")
}
func TestParseRateLimit(t *testing.T) {
tests := []struct {
input string
wantLen int
}{
{"10/sec", 1},
{"5/min", 1},
{"100/hour", 1},
{"1000/day", 1},
{"s:10/sec:20", 1},
{"invalid", 0},
{"", 0},
}
for _, tt := range tests {
exprs := parseRateLimit(tt.input)
if len(exprs) != tt.wantLen {
t.Errorf("parseRateLimit(%q) returned %d exprs, want %d", tt.input, len(exprs), tt.wantLen)
}
}
}
func TestCompile_RateLimit(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{
{
Action: config.RuleAccept,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"22"},
RateLimit: "10/sec:5",
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
hasLimit := false
for _, e := range r.Exprs {
if _, ok := e.(*expr.Limit); ok {
hasLimit = true
}
}
if !hasLimit {
t.Error("rule with rate_limit should have Limit expression")
}
return
}
}
t.Error("rule:0 not found in input chain")
}
func TestNegatedAddress(t *testing.T) {
exprs, err := matchSourceCIDR("!192.168.1.0/24")
if err != nil {
t.Fatalf("matchSourceCIDR(!192.168.1.0/24) error: %v", err)
}
if len(exprs) != 3 {
t.Fatalf("expected 3 exprs, got %d", len(exprs))
}
cmp := exprs[2].(*expr.Cmp)
if cmp.Op != expr.CmpOpNeq {
t.Errorf("negated address should use CmpOpNeq, got %v", cmp.Op)
}
exprs, err = matchDestCIDR("!10.0.0.1")
if err != nil {
t.Fatalf("matchDestCIDR(!10.0.0.1) error: %v", err)
}
if len(exprs) != 2 {
t.Fatalf("expected 2 exprs, got %d", len(exprs))
}
cmp = exprs[1].(*expr.Cmp)
if cmp.Op != expr.CmpOpNeq {
t.Errorf("negated address should use CmpOpNeq, got %v", cmp.Op)
}
}
func TestRejectTCPRST(t *testing.T) {
exprs := rejectExprs(unix.IPPROTO_TCP, config.FamilyINET)
if len(exprs) != 1 {
t.Fatalf("expected 1 expr, got %d", len(exprs))
}
rej := exprs[0].(*expr.Reject)
if rej.Type != 1 {
t.Errorf("TCP reject should use NFT_REJECT_TCP_RST (1), got %d", rej.Type)
}
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)
}
}
func TestCompile_DHCP(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", Options: config.InterfaceOptions{DHCP: true}},
},
Policy: []config.Policy{
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
foundIn := false
foundOut := false
for _, r := range state.Rules["input"] {
if r.Tag == "dhcp:in:eth0" || r.Tag == "dhcp:reply:eth0" {
foundIn = true
}
}
for _, r := range state.Rules["output"] {
if r.Tag == "dhcp:out:eth0" {
foundOut = true
}
}
if !foundIn {
t.Error("no DHCP input rule found for eth0")
}
if !foundOut {
t.Error("no DHCP output rule found for eth0")
}
}
func TestMatchUID(t *testing.T) {
exprs := matchUID("1000")
if len(exprs) != 2 {
t.Fatalf("matchUID(1000) returned %d exprs, want 2", len(exprs))
}
cmp := exprs[1].(*expr.Cmp)
if cmp.Op != expr.CmpOpEq {
t.Error("non-negated UID should use CmpOpEq")
}
exprs = matchUID("!0")
if len(exprs) != 2 {
t.Fatalf("matchUID(!0) returned %d exprs, want 2", len(exprs))
}
cmp = exprs[1].(*expr.Cmp)
if cmp.Op != expr.CmpOpNeq {
t.Error("negated UID should use CmpOpNeq")
}
}
func TestMatchMark(t *testing.T) {
exprs := matchMark("0x10/0xff")
if len(exprs) != 3 {
t.Fatalf("matchMark(0x10/0xff) returned %d exprs, want 3 (load+bitwise+cmp)", len(exprs))
}
exprs = matchMark("42")
if len(exprs) != 2 {
t.Fatalf("matchMark(42) returned %d exprs, want 2 (load+cmp)", len(exprs))
}
exprs = matchMark("!5")
cmp := exprs[1].(*expr.Cmp)
if cmp.Op != expr.CmpOpNeq {
t.Error("negated mark should use CmpOpNeq")
}
}
func TestSetMarkExprs(t *testing.T) {
exprs := setMarkExprs("0x10")
if len(exprs) != 2 {
t.Fatalf("setMarkExprs(0x10) returned %d exprs, want 2", len(exprs))
}
exprs = setMarkExprs("0x10/0xff00")
if len(exprs) != 3 {
t.Fatalf("setMarkExprs(0x10/0xff00) returned %d exprs, want 3", len(exprs))
}
}
func TestCompile_Loopback(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},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
foundInput := false
foundOutput := false
for _, r := range state.Rules["input"] {
if r.Tag == "loopback:input" {
foundInput = true
}
}
for _, r := range state.Rules["output"] {
if r.Tag == "loopback:output" {
foundOutput = true
}
}
if !foundInput {
t.Error("loopback:input rule not found")
}
if !foundOutput {
t.Error("loopback:output rule not found")
}
}
func TestCompile_AntiSpoof(t *testing.T) {
tcpflags := true
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", Options: config.InterfaceOptions{
NoSmurfs: true,
TCPFlags: &tcpflags,
}},
},
Policy: []config.Policy{
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
foundSmurf := false
foundFlags := false
for _, r := range state.Rules["input"] {
if r.Tag == "antismurf:eth0" {
foundSmurf = true
}
if r.Tag == "tcpflags:eth0" {
foundFlags = true
}
}
if !foundSmurf {
t.Error("antismurf rule not found")
}
if !foundFlags {
t.Error("tcpflags rule not found")
}
}
func TestCompile_MSSClamp(t *testing.T) {
cfg := &config.Config{
Settings: config.Settings{
TableName: "test",
AddressFamily: config.FamilyINET,
},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"loc": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "loc", Interface: "eth1", Options: config.InterfaceOptions{MSS: 1400}},
},
Policy: []config.Policy{
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
var rule *ManagedRule
for i, r := range state.Rules["forward"] {
if r.Tag == "mss:eth1" {
rule = &state.Rules["forward"][i]
}
}
if rule == nil {
t.Fatal("MSS clamp rule not found in forward chain")
}
mss := []byte{0x05, 0x78}
want := []expr.Any{
&expr.Exthdr{DestRegister: 1, Type: 2, Offset: 2, Len: 2, Op: expr.ExthdrOpTcpopt},
&expr.Cmp{Op: expr.CmpOpGt, Register: 1, Data: mss},
&expr.Immediate{Register: 1, Data: mss},
&expr.Exthdr{SourceRegister: 1, Type: 2, Offset: 2, Len: 2, Op: expr.ExthdrOpTcpopt},
}
got := rule.Exprs[len(rule.Exprs)-len(want):]
if !reflect.DeepEqual(got, want) {
t.Errorf("MSS clamp exprs = %#v, want %#v", got, want)
}
}
func TestCompile_PolicyExclusion(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},
"loc": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "loc", Interface: "eth1"},
},
Policy: []config.Policy{
{Source: "all!net", Dest: "all", Action: config.PolicyAccept},
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["forward"] {
if r.Tag == "policy:0" {
return
}
}
for _, r := range state.Rules["input"] {
if r.Tag == "policy:0" {
return
}
}
for _, r := range state.Rules["output"] {
if r.Tag == "policy:0" {
return
}
}
t.Error("policy:0 (all!net exclusion) not found in any chain")
}
func TestMatchICMPType(t *testing.T) {
tests := []struct {
input string
wantLen int
}{
{"echo-request", 2},
{"8", 2},
{"3/4", 4},
{"destination-unreachable", 2},
}
for _, tt := range tests {
exprs, err := matchICMPType(tt.input)
if err != nil {
t.Fatalf("matchICMPType(%q) error: %v", tt.input, err)
}
if len(exprs) != tt.wantLen {
t.Errorf("matchICMPType(%q) returned %d exprs, want %d", tt.input, len(exprs), tt.wantLen)
}
}
}
func TestCompile_ICMPRule(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{
{
Action: config.RuleAccept,
Source: "net",
Dest: "fw",
Proto: "icmp",
DPort: config.PortSpec{"echo-request"},
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
if len(r.Exprs) < 5 {
t.Errorf("ICMP rule should have >= 5 exprs (iface+proto+icmptype+verdict), got %d", len(r.Exprs))
}
return
}
}
t.Error("rule:0 not found in input chain")
}
func TestMatchConnLimit(t *testing.T) {
exprs := matchConnLimit("20")
if len(exprs) != 1 {
t.Fatalf("matchConnLimit(20) returned %d exprs, want 1", len(exprs))
}
cl := exprs[0].(*expr.Connlimit)
if cl.Count != 20 {
t.Errorf("Connlimit.Count = %d, want 20", cl.Count)
}
exprs = matchConnLimit("d:10")
cl = exprs[0].(*expr.Connlimit)
if cl.Flags != 1 {
t.Errorf("d: prefix should set Flags=1, got %d", cl.Flags)
}
}
func TestCompile_ConnLimit(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{
{
Action: config.RuleAccept,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"22"},
ConnLimit: "20",
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
hasConnLimit := false
for _, e := range r.Exprs {
if _, ok := e.(*expr.Connlimit); ok {
hasConnLimit = true
}
}
if !hasConnLimit {
t.Error("rule with conn_limit should have Connlimit expression")
}
return
}
}
t.Error("rule:0 not found in input chain")
}
func TestCompile_NFQUEUE(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{
{
Action: config.RuleNFQueue,
Source: "net",
Dest: "fw",
Proto: "tcp",
DPort: config.PortSpec{"80"},
NFQueue: 1,
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["input"] {
if r.Tag == "rule:0" {
hasQueue := false
for _, e := range r.Exprs {
if q, ok := e.(*expr.Queue); ok {
hasQueue = true
if q.Num != 1 {
t.Errorf("Queue.Num = %d, want 1", q.Num)
}
}
}
if !hasQueue {
t.Error("NFQUEUE rule should have Queue expression")
}
return
}
}
t.Error("rule:0 not found in input chain")
}
func TestCompile_NONAT(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},
"loc": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "loc", Interface: "eth1"},
},
Policy: []config.Policy{
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
Rules: []config.Rule{
{
Action: config.RuleNoNAT,
Source: "net",
Dest: "loc",
Proto: "tcp",
DPort: config.PortSpec{"80"},
},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["forward"] {
if r.Tag == "rule:0" {
hasReturn := false
for _, e := range r.Exprs {
if v, ok := e.(*expr.Verdict); ok && v.Kind == expr.VerdictReturn {
hasReturn = true
}
}
if !hasReturn {
t.Error("NONAT rule should have RETURN verdict")
}
return
}
}
t.Error("rule:0 not found in forward chain")
}
func TestCompile_PolicyRateLimit(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: "net", Dest: "fw", Action: config.PolicyDrop, RateLimit: "5/sec"},
{Source: "all", Dest: "all", Action: config.PolicyDrop},
},
PortGroups: make(map[string]config.PortGroup),
}
c := NewCompiler(cfg)
state, err := c.Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
for _, r := range state.Rules["input"] {
if r.Tag == "policy:0" {
hasLimit := false
for _, e := range r.Exprs {
if _, ok := e.(*expr.Limit); ok {
hasLimit = true
}
}
if !hasLimit {
t.Error("policy with rate_limit should have Limit expression")
}
return
}
}
t.Error("policy:0 not found in input chain")
}
func diffTestConfig(port string) *config.Config {
return &config.Config{
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
Zones: map[string]config.Zone{
"fw": {Type: config.ZoneFirewall},
"net": {Type: config.ZoneIP},
"loc": {Type: config.ZoneIP},
"dmz": {Type: config.ZoneIP},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"},
{Zone: "loc", Interface: "eth1"},
{Zone: "dmz", Interface: "eth2"},
},
Policy: []config.Policy{
{Source: "net", Dest: "all", Action: config.PolicyDrop, Log: "info"},
{Source: "all", Dest: "all", Action: config.PolicyReject},
},
Rules: []config.Rule{
{Action: config.RuleAccept, Source: "net:192.0.2.0/24", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{port}},
{Action: config.RuleDNAT, Source: "net", Dest: "loc:198.51.100.10:80", Proto: "tcp", DPort: config.PortSpec{"8000"}},
{Action: config.RuleNFQueue, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"80"}, NFQueue: 3},
{Action: config.RuleRedirect, Source: "loc", Dest: "fw:192.0.2.1:3128", Proto: "tcp", DPort: config.PortSpec{"80"}},
{Action: config.RuleAccept, Source: "loc,dmz", Dest: "fw,net", Proto: "tcp,udp", DPort: config.PortSpec{"53", "5353"}},
},
SNAT: []config.SNATRule{{Action: config.SNATAddress, Source: "198.51.100.0/24", Dest: "eth0", Address: "203.0.113.7"}},
PortGroups: make(map[string]config.PortGroup),
}
}
func TestDiffEngine_IndependentCompilesMatch(t *testing.T) {
compile := func(port string) *FirewallState {
t.Helper()
s, err := NewCompiler(diffTestConfig(port)).Compile()
if err != nil {
t.Fatal(err)
}
return s
}
for i := 0; i < 20; i++ {
if cs := computeDiff(compile("22"), compile("22")); !cs.Empty() {
t.Fatalf("expected empty changeset, got:\n%s", cs.Summary())
}
}
cs := computeDiff(compile("22"), compile("2222"))
if len(cs.Add) != 1 || len(cs.Remove) != 1 || cs.Add[0].Tag != "rule:0" {
t.Errorf("expected rule:0 replaced, got:\n%s", cs.Summary())
}
}
// shapes observed by applying diffTestConfig in a netns and reading it back
func TestCompile_QueueRedirMatchKernelReadback(t *testing.T) {
state, err := NewCompiler(diffTestConfig("22")).Compile()
if err != nil {
t.Fatal(err)
}
want := map[string]expr.Any{
"rule:2": &expr.Queue{Num: 3, Total: 1},
"rule:3": &expr.Redir{RegisterProtoMin: 1, RegisterProtoMax: 1, Flags: unix.NF_NAT_RANGE_PROTO_SPECIFIED},
}
for _, rules := range state.Rules {
for _, r := range rules {
w, ok := want[r.Tag]
if !ok {
continue
}
if got := r.Exprs[len(r.Exprs)-1]; !reflect.DeepEqual(got, w) {
t.Errorf("%s: got %#v, want %#v", r.Tag, got, w)
}
delete(want, r.Tag)
}
}
for tag := range want {
t.Errorf("%s not compiled", tag)
}
}
func withHandles(s *FirewallState) *FirewallState {
h := uint64(100)
for _, rules := range s.Rules {
for i := range rules {
rules[i].Handle = h
h++
}
}
return s
}
func tags(rules []ManagedRule) []string {
out := make([]string, len(rules))
for i, r := range rules {
out[i] = r.Chain + "/" + r.Tag
}
return out
}
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"}},
},
{
name: "protocol names",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "ospf,OSPFIGP,igmp,gre,esp,ah,vrrp,pim,ipencap,ipv6-icmp"},
chain: "input",
want: []want{{89, ""}, {89, ""}, {2, ""}, {47, ""}, {50, ""}, {51, ""}, {112, ""}, {103, ""}, {4, ""}, {58, ""}},
},
{
name: "protocol numbers",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "0,89,255"},
chain: "input",
want: []want{{0, ""}, {89, ""}, {255, ""}},
},
{
name: "numeric tcp with port",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "6,udplite", DPort: config.PortSpec{"22"}},
chain: "input",
want: []want{{6, "22"}, {136, "22"}},
},
}
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)
}
})
}
}
func TestCompile_OutputPolicyMatchesOif(t *testing.T) {
cfg := listCfg(func(c *config.Config) {
c.Zones["lan"] = config.Zone{Type: config.ZoneIP}
c.Interfaces = append(c.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"})
c.Policy = []config.Policy{{Source: "fw", Dest: "lan", Action: config.PolicyAccept}}
})
got := taggedRules(mustCompile(t, cfg), "output", "policy:0")
if len(got) != 1 || describeRule(got[0]) != "oif=eth1" {
t.Fatalf("fw->lan policy = %v, want one rule oif=eth1", got)
}
}
func TestCompile_InterfacelessZonesFailClosed(t *testing.T) {
ipsec := func(c *config.Config) { c.Zones["ips"] = config.Zone{Type: config.ZoneIPSec} }
rule := func(dest string) func(*config.Config) {
return func(c *config.Config) {
ipsec(c)
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "fw", Dest: dest}}
}
}
tests := []struct {
name string
mod func(*config.Config)
tag string
want int
warns []string
}{
{"ipsec zone without interface", func(c *config.Config) {
ipsec(c)
c.Policy = []config.Policy{{Source: "fw", Dest: "ips", Action: config.PolicyAccept}}
}, "policy:0", 0, []string{"ips"}},
{"hosts-only zone", func(c *config.Config) {
c.Zones["hst"] = config.Zone{Type: config.ZoneIP}
c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}}
c.Policy = []config.Policy{{Source: "fw", Dest: "hst", Action: config.PolicyAccept}}
}, "policy:0", 0, []string{"hst"}},
{"fw all expansion keeps zones with interfaces", func(c *config.Config) {
ipsec(c)
c.Zones["hst"] = config.Zone{Type: config.ZoneIP}
c.Hosts = []config.Host{{Zone: "hst", Interface: "eth0", Addresses: []string{"192.0.2.0/24"}}}
c.Policy = []config.Policy{{Source: "fw", Dest: "all", Action: config.PolicyDrop}}
}, "policy:0", 1, []string{"hst", "ips"}},
{"negated address does not scope", rule("ips:!192.0.2.1"), "rule:0", 0, []string{"ips"}},
{"address scopes", rule("ips:192.0.2.1"), "rule:0", 1, []string{"ips"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var logs bytes.Buffer
prev := slog.Default()
slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil)))
defer slog.SetDefault(prev)
state := mustCompile(t, listCfg(tt.mod))
got := 0
for chain := range state.Rules {
for _, r := range taggedRules(state, chain, tt.tag) {
got++
if describeRule(r) == "" {
t.Errorf("%s: un-scoped rule", chain)
}
}
}
if got != tt.want {
t.Errorf("got %d %s rules, want %d", got, tt.tag, tt.want)
}
if n := strings.Count(logs.String(), "zone has no interfaces"); n != len(tt.warns) {
t.Errorf("got %d warnings, want %d:\n%s", n, len(tt.warns), logs.String())
}
for _, z := range tt.warns {
if n := strings.Count(logs.String(), "zone="+z+"\n"); n != 1 {
t.Errorf("zone %s warned %d times, want 1", z, n)
}
}
})
}
}
// describeRule renders a rule's iif/oif/saddr/daddr matches, e.g. "iif=eth1 oif=eth2 daddr=192.0.2.1".
func describeRule(r ManagedRule) string {
var parts []string
for i, e := range r.Exprs {
cmp, ok := func() (*expr.Cmp, bool) {
if i+1 >= len(r.Exprs) {
return nil, false
}
c, ok := r.Exprs[i+1].(*expr.Cmp)
return c, ok
}()
if !ok {
continue
}
switch m := e.(type) {
case *expr.Meta:
switch m.Key {
case expr.MetaKeyIIFNAME:
parts = append(parts, "iif="+strings.TrimRight(string(cmp.Data), "\x00"))
case expr.MetaKeyOIFNAME:
parts = append(parts, "oif="+strings.TrimRight(string(cmp.Data), "\x00"))
}
case *expr.Payload:
if m.Base == expr.PayloadBaseNetworkHeader && (m.Len == 4 || m.Len == 16) {
name := map[uint32]string{12: "saddr", 16: "daddr", 8: "saddr", 24: "daddr"}[m.Offset]
if cmp.Op == expr.CmpOpNeq {
name = "!" + name
}
parts = append(parts, name+"="+net.IP(cmp.Data).String())
}
}
}
return strings.Join(parts, " ")
}
func TestCompile_CommaZoneLists(t *testing.T) {
tests := []struct {
name string
rule config.Rule
blrule *config.BlruleRule
want map[string][]string
}{
{
name: "fw in source list goes to output",
rule: config.Rule{Action: config.RuleAccept, Source: "fw,lan", Dest: "svr", Proto: "tcp", DPort: config.PortSpec{"22"}},
want: map[string][]string{"output": {"oif=eth2"}, "forward": {"iif=eth1 oif=eth2"}},
},
{
name: "dest list with fw splits input and forward",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "fw,svr,net"},
want: map[string][]string{"input": {"iif=eth1"}, "forward": {"iif=eth1 oif=eth2", "iif=eth1 oif=eth0"}},
},
{
name: "zone without interfaces emits nothing",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "svr,dmz"},
want: map[string][]string{"forward": {"iif=eth1 oif=eth2"}},
},
{
name: "address list after colon belongs to one zone",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "net:192.0.2.1,198.51.100.1"},
want: map[string][]string{"forward": {"iif=eth1 oif=eth0 daddr=192.0.2.1", "iif=eth1 oif=eth0 daddr=198.51.100.1"}},
},
{
name: "zone:address inside a list",
rule: config.Rule{Action: config.RuleAccept, Source: "lan,svr:203.0.113.7", Dest: "fw"},
want: map[string][]string{"input": {"iif=eth1", "iif=eth2 saddr=203.0.113.7"}},
},
{
name: "dnat source list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth1"}},
},
{
name: "dnat source address list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net:192.0.2.5,198.51.100.5", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}},
want: map[string][]string{"prerouting": {"iif=eth0 saddr=192.0.2.5", "iif=eth0 saddr=198.51.100.5"}},
},
{
name: "negated address list stays one AND-ed rule",
rule: config.Rule{Action: config.RuleAccept, Source: "net:!192.0.2.5,198.51.100.5", Dest: "fw"},
want: map[string][]string{"input": {"iif=eth0 !saddr=192.0.2.5 !saddr=198.51.100.5"}},
},
{
name: "zone named like all/any keyword is a plain zone",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "anycast,net"},
want: map[string][]string{"forward": {"iif=eth1 oif=eth3", "iif=eth1 oif=eth0"}},
},
{
name: "interface-less zones are skipped",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn,dmz"},
want: map[string][]string{},
},
{
name: "interface-less zone kept when address narrows it",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn:192.0.2.1"},
want: map[string][]string{"forward": {"iif=eth1 daddr=192.0.2.1"}},
},
{
name: "fw source matches dest zone oif",
rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "lan,vpn,dmz", Proto: "tcp", DPort: config.PortSpec{"22"}},
want: map[string][]string{"output": {"oif=eth1"}},
},
{
name: "fw to all has no oif",
rule: config.Rule{Action: config.RuleAccept, Source: "fw", Dest: "all:192.0.2.1"},
want: map[string][]string{"output": {"daddr=192.0.2.1"}},
},
{
name: "dnat origdest",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"},
want: map[string][]string{"prerouting": {"iif=eth0 daddr=203.0.113.5"}},
},
{
name: "dnat origdest list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5,203.0.113.6"},
want: map[string][]string{"prerouting": {"iif=eth0 daddr=203.0.113.5", "iif=eth0 daddr=203.0.113.6"}},
},
{
name: "dnat negated origdest list",
rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "!203.0.113.5,203.0.113.6"},
want: map[string][]string{"prerouting": {"iif=eth0 !daddr=203.0.113.5 !daddr=203.0.113.6"}},
},
{
name: "accept origdest",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"},
want: map[string][]string{"input": {"iif=eth0 daddr=203.0.113.5"}},
},
{
name: "origdest does not scope interface-less zone",
rule: config.Rule{Action: config.RuleAccept, Source: "vpn", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "203.0.113.5"},
want: map[string][]string{},
},
{
name: "accept ipv6 origdest",
rule: config.Rule{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, OrigDest: "2001:db8::5"},
want: map[string][]string{"input": {"iif=eth0 daddr=2001:db8::5"}},
},
{
name: "blrule zone list",
blrule: &config.BlruleRule{Action: config.BlruleDrop, Source: "net,anycast", Dest: "fw"},
want: map[string][]string{"input": {"iif=eth0", "iif=eth3"}},
},
}
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},
"lan": {Type: config.ZoneIP}, "svr": {Type: config.ZoneIP}, "dmz": {Type: config.ZoneIP},
"anycast": {Type: config.ZoneIP}, "vpn": {Type: config.ZoneIPSec},
},
Interfaces: []config.Interface{
{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}, {Zone: "svr", Interface: "eth2"},
{Zone: "anycast", Interface: "eth3"},
},
Rules: []config.Rule{tt.rule},
PortGroups: make(map[string]config.PortGroup),
}
tag := "rule:0"
if tt.blrule != nil {
cfg.Rules, cfg.Blrules, tag = nil, []config.BlruleRule{*tt.blrule}, "blrule:0"
}
state, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
got := map[string][]string{}
for chain, rules := range state.Rules {
for _, r := range rules {
if r.Tag == tag {
got[chain] = append(got[chain], describeRule(r))
}
}
}
if !reflect.DeepEqual(got, tt.want) {
t.Errorf("rules = %v, want %v", got, tt.want)
}
})
}
}
func listCfg(mod func(*config.Config)) *config.Config {
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}},
PortGroups: make(map[string]config.PortGroup),
}
mod(cfg)
return cfg
}
func taggedRules(state *FirewallState, chain, tag string) []ManagedRule {
var out []ManagedRule
for _, r := range state.Rules[chain] {
if r.Tag == tag {
out = append(out, r)
}
}
return out
}
func TestDiffEngine_FreshApplyKeepsDesiredOrder(t *testing.T) {
desired, err := NewCompiler(diffTestConfig("22")).Compile()
if err != nil {
t.Fatal(err)
}
var want []string
for _, chain := range []string{"forward", "input", "output", "postrouting", "prerouting"} {
want = append(want, tags(desired.Rules[chain])...)
}
for i := 0; i < 20; i++ {
cs := computeDiff(&FirewallState{Rules: map[string][]ManagedRule{}}, desired)
if got := tags(cs.Add); !reflect.DeepEqual(got, want) {
t.Fatalf("add order:\n got %v\nwant %v", got, want)
}
for _, r := range cs.Add {
if r.Before != 0 {
t.Fatalf("fresh apply should append, %s has Before=%d", r.Tag, r.Before)
}
}
}
}
func TestDiffEngine_MiddleChangeInsertsBeforeNextRule(t *testing.T) {
current := withHandles(mustCompile(t, diffTestConfig("22")))
desired := mustCompile(t, diffTestConfig("2222"))
cs := computeDiff(current, desired)
if len(cs.Add) != 1 || len(cs.Remove) != 1 {
t.Fatalf("expected one replace, got:\n%s", cs.Summary())
}
input := current.Rules["input"]
idx := -1
for i, r := range input {
if r.Tag == "rule:0" {
idx = i
}
}
if idx < 1 || idx == len(input)-1 {
t.Fatalf("rule:0 at %d is not mid-chain", idx)
}
if cs.Remove[0].Handle != input[idx].Handle {
t.Errorf("removed handle %d, want %d", cs.Remove[0].Handle, input[idx].Handle)
}
if cs.Add[0].Before != input[idx+1].Handle {
t.Errorf("insert before %d, want %d (%s)", cs.Add[0].Before, input[idx+1].Handle, input[idx+1].Tag)
}
}
func TestDiffEngine_ExpandedRuleReplacedInPlace(t *testing.T) {
current := withHandles(mustCompile(t, diffTestConfig("22")))
if cs := computeDiff(current, mustCompile(t, diffTestConfig("22"))); !cs.Empty() {
t.Fatalf("expected empty changeset, got:\n%s", cs.Summary())
}
cfg := diffTestConfig("22")
cfg.Rules[4].DPort = config.PortSpec{"53", "853"}
desired := mustCompile(t, cfg)
cs := computeDiff(current, desired)
for _, r := range append(append([]ManagedRule{}, cs.Add...), cs.Remove...) {
if r.Tag != "rule:4" {
t.Errorf("unexpected change to %s/%s", r.Chain, r.Tag)
}
}
for _, chain := range []string{"input", "forward"} {
if n := len(taggedRules(desired, chain, "rule:4")); n != 8 {
t.Fatalf("%s: expected 8 expanded rule:4 rules, got %d", chain, n)
}
}
applied := applyChangeSet(current, cs)
if cs := computeDiff(applied, desired); !cs.Empty() {
t.Fatalf("second plan not empty:\n%s", cs.Summary())
}
}
// applyChangeSet mimics the engine: removals by handle, adds inserted before r.Before or appended.
func applyChangeSet(s *FirewallState, cs *ChangeSet) *FirewallState {
gone := map[uint64]bool{}
for _, r := range cs.Remove {
gone[r.Handle] = true
}
out := &FirewallState{Rules: map[string][]ManagedRule{}}
for chain, rules := range s.Rules {
for _, r := range rules {
if !gone[r.Handle] {
out.Rules[chain] = append(out.Rules[chain], r)
}
}
}
h := uint64(10000)
for _, r := range cs.Add {
r.Handle, h = h, h+1
rules := out.Rules[r.Chain]
i := len(rules)
for j, x := range rules {
if r.Before != 0 && x.Handle == r.Before {
i = j
break
}
}
r.Before = 0
out.Rules[r.Chain] = append(rules[:i], append([]ManagedRule{r}, rules[i:]...)...)
}
return out
}
func mustCompile(t *testing.T, cfg *config.Config) *FirewallState {
t.Helper()
s, err := NewCompiler(cfg).Compile()
if err != nil {
t.Fatal(err)
}
return s
}
func TestCompile_ListExpansionCounts(t *testing.T) {
tests := []struct {
name string
mod func(*config.Config)
chain string
tag string
want int
}{
{"proto x dport cross product", func(c *config.Config) {
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80", "443"}}}
}, "input", "rule:0", 4},
{"sport list", 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", 2},
{"snat proto x dport", func(c *config.Config) {
c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}}
}, "postrouting", "snat:0", 4},
{"conntrack dport list", func(c *config.Config) {
c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53", "123"}}}
}, "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)
}
if got := len(taggedRules(state, tt.chain, tt.tag)); got != tt.want {
t.Errorf("%s rules in %s = %d, want %d", tt.tag, tt.chain, got, tt.want)
}
})
}
}
func TestCompile_DNATGetsNoRuleExtras(t *testing.T) {
compile := func(mark string) []expr.Any {
state, err := NewCompiler(listCfg(func(c *config.Config) {
c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}, Mark: mark}}
})).Compile()
if err != nil {
t.Fatalf("Compile() error: %v", err)
}
return state.Rules["prerouting"][len(state.Rules["prerouting"])-1].Exprs
}
if plain, marked := compile(""), compile("0x1"); !reflect.DeepEqual(plain, marked) {
t.Errorf("DNAT prerouting rule changed by mark extra: %d exprs vs %d", len(plain), len(marked))
}
}
func TestCompile_CommaZoneListLimitErrors(t *testing.T) {
for _, r := range []config.Rule{
{Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"},
{Action: config.RuleAccept, Source: "net,lan", Dest: "fw", ConnLimit: "10"},
{Action: config.RuleAccept, Source: "net", Dest: "fw:192.0.2.1,198.51.100.1", RateLimit: "10/sec"},
} {
t.Run(r.Source+">"+r.Dest, 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}, "lan": {Type: config.ZoneIP}},
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}},
Rules: []config.Rule{r},
PortGroups: make(map[string]config.PortGroup),
}
if _, err := NewCompiler(cfg).Compile(); err == nil {
t.Fatal("Compile() succeeded, want error")
}
})
}
}
func TestCompile_RejectPerProto(t *testing.T) {
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}},
}
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])
}
}
})
}
}
func TestNegatedAddressList(t *testing.T) {
exprs, err := matchDestCIDR("!192.0.2.1,198.51.100.1")
if err != nil {
t.Fatalf("matchDestCIDR error: %v", err)
}
if len(exprs) != 4 {
t.Fatalf("expected 4 exprs, got %d", len(exprs))
}
for _, i := range []int{1, 3} {
if exprs[i].(*expr.Cmp).Op != expr.CmpOpNeq {
t.Errorf("expr %d should be CmpOpNeq", i)
}
}
}
func TestCompile_ListErrors(t *testing.T) {
tests := []struct {
name string
rule config.Rule
}{
{"invalid port in list", config.Rule{Proto: "tcp", DPort: config.PortSpec{"80,abc"}}},
{"trailing empty proto", config.Rule{Proto: "tcp,", DPort: config.PortSpec{"80"}}},
{"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"}},
{"unknown proto name", config.Rule{Proto: "bogus"}},
{"dport with ospf", config.Rule{Proto: "ospf", DPort: config.PortSpec{"80"}}},
{"sport with gre", config.Rule{Proto: "gre", SPort: config.PortSpec{"80"}}},
{"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"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
r := tt.rule
r.Action, r.Source, r.Dest = config.RuleAccept, "net", "fw"
_, err := NewCompiler(listCfg(func(c *config.Config) { c.Rules = []config.Rule{r} })).Compile()
if err == nil {
t.Fatal("Compile() succeeded, want error")
}
})
}
}
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)
}
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)
}
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)
}
})
}
}
func TestMatchOrigDest_FamilyGuard(t *testing.T) {
for addr, want := range map[string]byte{
"203.0.113.5": unix.NFPROTO_IPV4,
"!203.0.113.0/24,192.0.2.1": unix.NFPROTO_IPV4,
"2001:db8::5": unix.NFPROTO_IPV6,
} {
e, err := matchOrigDest(addr)
if err != nil {
t.Fatalf("%s: %v", addr, err)
}
if m, ok := e[0].(*expr.Meta); !ok || m.Key != expr.MetaKeyNFPROTO || e[1].(*expr.Cmp).Data[0] != want {
t.Errorf("%s: missing nfproto %d guard: %v", addr, want, e[:2])
}
}
if _, err := matchOrigDest("!203.0.113.5,2001:db8::5"); err == nil {
t.Error("mixed IPv4/IPv6 origdest: want error")
}
}
func TestCompile_OrigDestForwardRejected(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}, "svr": {Type: config.ZoneIP}},
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "svr", Interface: "eth2"}},
Rules: []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "svr", Proto: "tcp", DPort: config.PortSpec{"80"}, OrigDest: "203.0.113.5"}},
PortGroups: make(map[string]config.PortGroup),
}
_, err := NewCompiler(cfg).Compile()
if err == nil || !strings.Contains(err.Error(), "not supported yet") {
t.Fatalf("Compile() error = %v, want forwarded ORIGDEST rejection", err)
}
}