Files
tomswall/internal/nftables/compiler_test.go
T
unkin-agent abbf434af5
ci/woodpecker/pr/test Pipeline was successful
ci/woodpecker/pr/build Pipeline was successful
ci/woodpecker/pr/pre-commit Pipeline was successful
Merge remote-tracking branch 'origin/main' into benvin/mss-clamp-tcpopt
# Conflicts:
#	internal/nftables/compiler_test.go
2026-10-03 21:14:10 +10:00

2045 lines
53 KiB
Go

package nftables
import (
"encoding/binary"
"fmt"
"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)
}
found := false
for _, r := range state.Rules["prerouting"] {
if r.Tag == "rule:0" {
found = true
break
}
}
if !found {
t.Error("no DNAT rule found in prerouting chain")
}
}
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 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)
}
})
}
}
// 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 {
name := map[uint32]string{12: "saddr", 16: "daddr"}[m.Offset]
if cmp.Op == expr.CmpOpNeq {
name = "!" + name
}
parts = append(parts, fmt.Sprintf("%s=%d.%d.%d.%d", name, cmp.Data[0], cmp.Data[1], cmp.Data[2], cmp.Data[3]))
}
}
}
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": {""}, "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 ipsec zone keeps zone-agnostic rule, ip zone skipped",
rule: config.Rule{Action: config.RuleAccept, Source: "lan", Dest: "vpn,dmz"},
want: map[string][]string{"forward": {"iif=eth1"}},
},
{
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 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)
}
})
}
}