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:192.0.2.1", 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["raw_prerouting"] { if r.Tag == "conntrack:0:raw_prerouting" { found = true break } } if !found { t.Error("no notrack rule found in raw_prerouting chain") } if len(state.Rules["prerouting"]) != 0 { t.Error("conntrack rule leaked into the nat 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()) } if m.Base == expr.PayloadBaseTransportHeader && m.Offset == 0 && m.Len == 2 && cmp.Op == expr.CmpOpEq { parts = append(parts, fmt.Sprintf("sport=%d", binary.BigEndian.Uint16(cmp.Data))) } } } 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"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat source list skips the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "net,svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth0"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat lone source zone may equal the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "svr", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth2"}, "forward": {"iif=eth2 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat exclusion source skips the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "all!fw,anycast", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth1", "iif=eth0"}, "forward": {"iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat intrazone exclusion source keeps the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "all+!fw,anycast,lan", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth0", "iif=eth2"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth2 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat all source expands per zone and skips fw and the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "all", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0"}, "forward": {"iif=eth3 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat any+ source keeps the target zone", rule: config.Rule{Action: config.RuleDNAT, Source: "any+", Dest: "svr:192.0.2.10", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth3", "iif=eth1", "iif=eth0", "iif=eth2"}, "forward": {"iif=eth3 oif=eth2 daddr=192.0.2.10", "iif=eth1 oif=eth2 daddr=192.0.2.10", "iif=eth0 oif=eth2 daddr=192.0.2.10", "iif=eth2 oif=eth2 daddr=192.0.2.10"}}, }, { name: "dnat to fw accepts in input", rule: config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "fw:192.0.2.1", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth0"}, "input": {"iif=eth0 daddr=192.0.2.1"}}, }, { name: "redirect accepts in input without daddr", rule: config.Rule{Action: config.RuleRedirect, Source: "lan", Dest: "fw:192.0.2.1:3128", Proto: "tcp", DPort: config.PortSpec{"80"}}, want: map[string][]string{"prerouting": {"iif=eth1"}, "input": {"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"}, "forward": {"iif=eth0 oif=eth2 saddr=192.0.2.5 daddr=192.0.2.10", "iif=eth0 oif=eth2 saddr=198.51.100.5 daddr=192.0.2.10"}}, }, { 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"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, }, { 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"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, }, { 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"}, "forward": {"iif=eth0 oif=eth2 daddr=192.0.2.17"}}, }, { 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 || r.Tag == tag+":accept" { 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"}}} }, "raw_prerouting", "conntrack:0:raw_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_DNATImpliedAccept(t *testing.T) { state, err := NewCompiler(listCfg(func(c *config.Config) { c.Zones["svr"] = config.Zone{Type: config.ZoneIP} c.Interfaces = append(c.Interfaces, config.Interface{Zone: "svr", Interface: "eth2"}) c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17:8080", Proto: "tcp", DPort: config.PortSpec{"80"}}} })).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) } fwd := taggedRules(state, "forward", "rule:0:accept") if len(fwd) != 1 { t.Fatalf("got %d forward accepts, want 1", len(fwd)) } want := append(append(append(append(append(matchIfaceName(true, "eth0"), matchIfaceName(false, "eth2")...), mustExprs(t)(matchDestCIDR("192.0.2.17"))...), mustExprs(t)(l4Exprs("tcp", "8080"))...), dnatStatusExprs...), &expr.Verdict{Kind: expr.VerdictAccept}) if !reflect.DeepEqual(fwd[0].Exprs, want) { t.Errorf("forward accept = %#v, want %#v", fwd[0].Exprs, want) } } // IPS_DST_NAT = 1<<5, hard-coded so a wrong ctStatusDNAT or ct key fails here. var dnatStatusExprs = []expr.Any{ &expr.Ct{Key: expr.CtKeySTATUS, Register: 1}, &expr.Bitwise{SourceRegister: 1, DestRegister: 1, Len: 4, Mask: binary.NativeEndian.AppendUint32(nil, 32), Xor: []byte{0, 0, 0, 0}}, &expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: []byte{0, 0, 0, 0}}, } func TestCompile_DNATImpliedAcceptGetsNoRuleExtras(t *testing.T) { compile := func(r config.Rule) *FirewallState { state, err := NewCompiler(listCfg(func(c *config.Config) { c.Zones["svr"] = config.Zone{Type: config.ZoneIP} c.Interfaces = append(c.Interfaces, config.Interface{Zone: "svr", Interface: "eth2"}) c.Rules = []config.Rule{r} })).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) } return state } plain := config.Rule{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}} extras := plain extras.RateLimit, extras.Mark, extras.User = "10/sec:5", "0x1", "root" want, got := compile(plain), compile(extras) if len(taggedRules(got, "forward", "rule:0:accept")) != 1 { t.Fatalf("want one forward accept, got %v", got.Rules["forward"]) } if !reflect.DeepEqual(got.Rules, want.Rules) { t.Errorf("ratelimit/mark/user changed the DNAT rules:\ngot %#v\nwant %#v", got.Rules, want.Rules) } } func TestCompile_DNATMatchesSport(t *testing.T) { state, err := NewCompiler(listCfg(func(c *config.Config) { c.Zones["svr"] = config.Zone{Type: config.ZoneIP} c.Interfaces = append(c.Interfaces, config.Interface{Zone: "svr", Interface: "eth2"}) c.Rules = []config.Rule{{Action: config.RuleDNAT, Source: "net", Dest: "svr:192.0.2.17", Proto: "tcp", DPort: config.PortSpec{"80"}, SPort: config.PortSpec{"1024"}}} })).Compile() if err != nil { t.Fatalf("Compile() error: %v", err) } m, err := l4Matches("tcp", config.PortSpec{"80"}, config.PortSpec{"1024"}) if err != nil { t.Fatal(err) } for _, r := range append(taggedRules(state, "prerouting", "rule:0"), taggedRules(state, "forward", "rule:0:accept")...) { if !containsExprs(r.Exprs, m[0].exprs) { t.Errorf("%s rule lacks the sport match: %#v", r.Chain, r.Exprs) } } } func containsExprs(haystack, needle []expr.Any) bool { for i := 0; i+len(needle) <= len(haystack); i++ { if reflect.DeepEqual(haystack[i:i+len(needle)], needle) { return true } } return false } func mustExprs(t *testing.T) func([]expr.Any, error) []expr.Any { return func(e []expr.Any, err error) []expr.Any { t.Helper() if err != nil { t.Fatal(err) } return e } } func l4Exprs(proto, port string) ([]expr.Any, error) { m, err := l4Matches(proto, config.PortSpec{port}, nil) if err != nil { return nil, err } return m[0].exprs, nil } 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"}, {Action: config.RuleDNAT, Source: "net,lan", Dest: "fw:192.0.2.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"}}} }, "raw_prerouting", "conntrack:0:raw_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) } } func TestCompile_ConntrackHelper(t *testing.T) { compile := func(rules ...config.ConntrackRule) (*FirewallState, error) { cfg := &config.Config{ Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET}, Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}}, Conntrack: rules, PortGroups: map[string]config.PortGroup{}, } return NewCompiler(cfg).Compile() } state, err := compile( config.ConntrackRule{Action: config.ConntrackHelper, Helper: "ftp", Proto: "tcp", DPort: config.PortSpec{"21"}}, config.ConntrackRule{Action: config.ConntrackHelper, Helper: "sip", Proto: "tcp", DPort: config.PortSpec{"5060"}, Chain: config.ConntrackPrerouting}, config.ConntrackRule{Action: config.ConntrackHelper, Helper: "tftp", Chain: config.ConntrackOutput}, ) if err != nil { t.Fatal(err) } want := []Helper{ {"ftp", expr.CtHelper{Name: "ftp", L3Proto: unix.NFPROTO_INET, L4Proto: unix.IPPROTO_TCP}}, {"sip-tcp", expr.CtHelper{Name: "sip", L3Proto: unix.NFPROTO_INET, L4Proto: unix.IPPROTO_TCP}}, {"tftp", expr.CtHelper{Name: "tftp", L3Proto: unix.NFPROTO_INET, L4Proto: unix.IPPROTO_UDP}}, } if !reflect.DeepEqual(state.Helpers, want) { t.Errorf("helpers = %+v, want %+v", state.Helpers, want) } refs := func(chain string) []string { var out []string for _, r := range state.Rules[chain] { ref, ok := r.Exprs[len(r.Exprs)-1].(*expr.Objref) if !ok || ref.Type != unix.NFT_OBJECT_CT_HELPER { t.Fatalf("%s: rule %s does not end in a ct helper objref", chain, r.Tag) } out = append(out, r.Tag+"="+ref.Name) } return out } if got := refs("helper_prerouting"); !reflect.DeepEqual(got, []string{ "conntrack:0:helper_prerouting=ftp", "conntrack:1:helper_prerouting=sip-tcp"}) { t.Errorf("helper_prerouting = %v", got) } if got := refs("helper_output"); !reflect.DeepEqual(got, []string{ "conntrack:0:helper_output=ftp", "conntrack:2:helper_output=tftp"}) { t.Errorf("helper_output = %v", got) } ftp := state.Rules["helper_prerouting"][0].Exprs l4, _ := l4Matches("tcp", config.PortSpec{"21"}, nil) if !reflect.DeepEqual(ftp[:len(ftp)-1], l4[0].exprs) { t.Errorf("ftp rule does not match tcp dport 21: %#v", ftp) } if _, err := compile(config.ConntrackRule{Action: config.ConntrackHelper, Helper: "nope"}); err == nil { t.Error("expected error for unknown helper without proto") } } func TestCompile_ConntrackZones(t *testing.T) { tests := []struct { name string ct config.ConntrackRule want map[string][]string wantErr string }{ { name: "source zone matches iif", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Proto: "udp", DPort: config.PortSpec{"53"}}, want: map[string][]string{"raw_prerouting": {"iif=eth0"}}, }, { name: "source and dest addresses", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:192.0.2.1,198.51.100.1", Dest: "fw:203.0.113.1"}, want: map[string][]string{"raw_prerouting": {"iif=eth0 saddr=192.0.2.1 daddr=203.0.113.1", "iif=eth0 saddr=198.51.100.1 daddr=203.0.113.1"}}, }, { name: "fw source goes to raw_output with dest oif", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "net,lan"}, want: map[string][]string{"raw_output": {"oif=eth0", "oif=eth1"}}, }, { name: "all matches no interface", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Dest: "fw:192.0.2.53"}, want: map[string][]string{"raw_prerouting": {"daddr=192.0.2.53"}}, }, { name: "interface-less zone fails closed", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "dmz"}, want: map[string][]string{}, }, { name: "chain output with fw source matches dest oif", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "net", Chain: config.ConntrackOutput}, want: map[string][]string{"raw_output": {"oif=eth0"}}, }, { name: "chain output with non-fw source is rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw", Chain: config.ConntrackOutput}, wantErr: "chain output needs SOURCE fw", }, { name: "chain prerouting with fw source is rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "net", Chain: config.ConntrackPrerouting}, wantErr: "SOURCE fw cannot use chain prerouting", }, { name: "omitted source and dest is global in both chains", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, want: map[string][]string{"raw_prerouting": {""}, "raw_output": {""}}, }, { name: "all source with chain both is global in both chains", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, want: map[string][]string{"raw_prerouting": {""}, "raw_output": {""}}, }, { name: "all source with chain output is global in raw_output", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackOutput}, want: map[string][]string{"raw_output": {""}}, }, { name: "any source with chain both is global in both chains", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "any", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, want: map[string][]string{"raw_prerouting": {""}, "raw_output": {""}}, }, { name: "any source with chain output is global in raw_output", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "any", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackOutput}, want: map[string][]string{"raw_output": {""}}, }, { name: "chain both with non-fw source emits prerouting only", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Proto: "udp", DPort: config.PortSpec{"53"}, Chain: config.ConntrackBoth}, want: map[string][]string{"raw_prerouting": {"iif=eth0"}}, }, { name: "chain both with fw source emits output only", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "fw", Dest: "lan", Chain: config.ConntrackBoth}, want: map[string][]string{"raw_output": {"oif=eth1"}}, }, { name: "negated addresses stay one AND-ed match", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net:!192.0.2.1,198.51.100.1"}, want: map[string][]string{"raw_prerouting": {"iif=eth0 !saddr=192.0.2.1 !saddr=198.51.100.1"}}, }, { name: "sport", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Proto: "udp", SPort: config.PortSpec{"123"}}, want: map[string][]string{"raw_prerouting": {"iif=eth0 sport=123"}}, }, { name: "dest zone without address is rejected in prerouting", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan"}, wantErr: `conntrack DEST zone "lan" needs an address in prerouting`, }, { name: "fw dest zone without address is rejected in prerouting", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net", Dest: "fw"}, wantErr: `conntrack DEST zone "fw" needs an address in prerouting`, }, { name: "fw dest zone with address matches daddr", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "fw:192.0.2.1"}, want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=192.0.2.1"}}, }, { name: "unknown zone fails closed", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "nte"}, want: map[string][]string{}, }, { name: "unknown dest zone fails closed in prerouting", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net", Dest: "typo"}, want: map[string][]string{}, }, { name: "none dest zone yields no rule", ct: config.ConntrackRule{Action: config.ConntrackDrop, Source: "net", Dest: "none"}, want: map[string][]string{}, }, { name: "all!net rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all!net"}, wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "all!net rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Dest: "all!net"}, wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "all+ rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "all+"}, wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "all+!net rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Dest: "all+!net"}, wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "any!net rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "any!net"}, wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "any+ rejected", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Dest: "any+"}, wantErr: "zone exclusions are not supported in conntrack entries", }, { name: "omitted source and dest is global", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53"}}, want: map[string][]string{"raw_prerouting": {""}}, }, { name: "dest zone with address matches daddr", ct: config.ConntrackRule{Action: config.ConntrackNoTrack, Source: "net", Dest: "lan:203.0.113.10"}, want: map[string][]string{"raw_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { cfg := listCfg(func(cfg *config.Config) { cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} cfg.Zones["dmz"] = config.Zone{Type: config.ZoneIP} cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) cfg.Conntrack = []config.ConntrackRule{tt.ct} }) if tt.wantErr != "" { if _, err := NewCompiler(cfg).Compile(); err == nil || !strings.Contains(err.Error(), tt.wantErr) { t.Fatalf("Compile() error = %v, want %q", err, tt.wantErr) } return } state := mustCompile(t, cfg) got := map[string][]string{} for _, chain := range []string{"raw_prerouting", "raw_output"} { for _, r := range taggedRules(state, chain, "conntrack:0:"+chain) { got[chain] = append(got[chain], describeRule(r)) } } if !reflect.DeepEqual(got, tt.want) { t.Errorf("rules = %v, want %v", got, tt.want) } }) } } func TestCompile_RuleZoneExclusionExpands(t *testing.T) { cfg := listCfg(func(cfg *config.Config) { cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) cfg.Rules = []config.Rule{ {Source: "all!net", Dest: "fw", Action: config.RuleAccept, Proto: "tcp", DPort: config.PortSpec{"22"}}, {Source: "nte", Dest: "fw", Action: config.RuleAccept}, } }) state := mustCompile(t, cfg) var got []string for _, r := range taggedRules(state, "input", "rule:0") { got = append(got, describeRule(r)) } if want := []string{"iif=eth1"}; !reflect.DeepEqual(got, want) { t.Errorf("all!net -> fw input rules = %v, want %v", got, want) } if r := taggedRules(state, "input", "rule:1"); len(r) != 0 { t.Errorf("unknown zone compiled %d rules, want 0", len(r)) } } func TestCompile_RuleZoneExclusionIntraZoneSymmetric(t *testing.T) { cfg := listCfg(func(cfg *config.Config) { cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) cfg.Rules = []config.Rule{ {Source: "lan", Dest: "all+!net", Action: config.RuleAccept}, {Source: "lan", Dest: "all!net", Action: config.RuleAccept}, } }) state := mustCompile(t, cfg) for i, want := range []bool{true, false} { got := false for _, r := range taggedRules(state, "forward", fmt.Sprintf("rule:%d", i)) { got = got || describeRule(r) == "iif=eth1 oif=eth1" } if got != want { t.Errorf("rule:%d lan->lan forward = %v, want %v", i, got, want) } } } func TestCompile_ConntrackHelperZones(t *testing.T) { tests := []struct { name string ct config.ConntrackRule want map[string][]string wantErr string }{ { name: "non-fw source is prerouting only", ct: config.ConntrackRule{Source: "net", Proto: "tcp", DPort: config.PortSpec{"21"}}, want: map[string][]string{"helper_prerouting": {"iif=eth0"}}, }, { name: "fw source is output only with dest oif", ct: config.ConntrackRule{Source: "fw", Dest: "lan", Proto: "tcp", DPort: config.PortSpec{"21"}}, want: map[string][]string{"helper_output": {"oif=eth1"}}, }, { name: "any source is global in both chains", ct: config.ConntrackRule{Source: "any", Proto: "tcp", DPort: config.PortSpec{"21"}}, want: map[string][]string{"helper_prerouting": {""}, "helper_output": {""}}, }, { name: "dest zone with address in prerouting", ct: config.ConntrackRule{Source: "net", Dest: "lan:203.0.113.10", Proto: "tcp", DPort: config.PortSpec{"21"}}, want: map[string][]string{"helper_prerouting": {"iif=eth0 daddr=203.0.113.10"}}, }, { name: "dest zone without address is rejected in prerouting", ct: config.ConntrackRule{Source: "net", Dest: "lan"}, wantErr: `conntrack DEST zone "lan" needs an address in prerouting`, }, { name: "chain output with non-fw source is rejected", ct: config.ConntrackRule{Source: "net", Chain: config.ConntrackOutput}, wantErr: "chain output needs SOURCE fw", }, { name: "interface-less zone fails closed", ct: config.ConntrackRule{Source: "dmz"}, want: map[string][]string{}, }, { name: "unknown zone fails closed", ct: config.ConntrackRule{Source: "nte"}, want: map[string][]string{}, }, { name: "exclusion rejected", ct: config.ConntrackRule{Source: "all!net"}, wantErr: "zone exclusions are not supported in conntrack entries", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { tt.ct.Action, tt.ct.Helper = config.ConntrackHelper, "ftp" cfg := listCfg(func(cfg *config.Config) { cfg.Zones["lan"] = config.Zone{Type: config.ZoneIP} cfg.Zones["dmz"] = config.Zone{Type: config.ZoneIP} cfg.Interfaces = append(cfg.Interfaces, config.Interface{Zone: "lan", Interface: "eth1"}) cfg.Conntrack = []config.ConntrackRule{tt.ct} }) if tt.wantErr != "" { if _, err := NewCompiler(cfg).Compile(); err == nil || !strings.Contains(err.Error(), tt.wantErr) { t.Fatalf("Compile() error = %v, want %q", err, tt.wantErr) } return } state := mustCompile(t, cfg) got := map[string][]string{} for _, chain := range []string{"helper_prerouting", "helper_output"} { for _, r := range taggedRules(state, chain, "conntrack:0:"+chain) { got[chain] = append(got[chain], describeRule(r)) } } if !reflect.DeepEqual(got, tt.want) { t.Errorf("rules = %v, want %v", got, tt.want) } }) } }