Files
tomswall/internal/config/extras_test.go
T
unkinben 8d9a76c751 Add comprehensive nftables compiler with shorewall feature parity
Rewrites the compiler from ~440 to ~1700 lines covering all major shorewall
firewall features: loopback, conntrack fast-path, anti-spoof, DHCP, intra-zone,
blacklist/whitelist, conntrack notrack, tunnels (13 types), rules with sections,
DNAT/redirect, SNAT/masquerade, static NAT, policies with zone exclusions,
MSS clamping, rate limiting, connection limiting, negated addresses, ICMP type
matching, TCP RST reject, user/UID matching, mark match/set, NFQUEUE, NONAT,
and policy-level rate/conn limiting.

Adds full config types for all shorewall subsystems (mangle, accounting, maclist,
netmap, providers, tunnels, conntrack, blrules, proxyarp/ndp, routes, tc, secmarks),
shorewall migration tooling, expanded CLI commands, expression-level diff engine,
and 49 unit tests.
2026-07-01 23:56:44 +10:00

1445 lines
30 KiB
Go

package config
import (
"strings"
"testing"
)
// --- 1. Conntrack ---
func TestValidateConntrack(t *testing.T) {
tests := []struct {
name string
rules []ConntrackRule
wantErr string
}{
{
name: "valid notrack rule",
rules: []ConntrackRule{
{Action: ConntrackNoTrack, Source: "net", Proto: "udp"},
},
},
{
name: "missing action",
rules: []ConntrackRule{
{Source: "net"},
},
wantErr: "unknown action",
},
{
name: "helper requires helper name",
rules: []ConntrackRule{
{Action: ConntrackHelper, Source: "net"},
},
wantErr: "helper name required",
},
{
name: "valid helper with name",
rules: []ConntrackRule{
{Action: ConntrackHelper, Source: "net", Helper: "ftp"},
},
},
{
name: "user requires output chain",
rules: []ConntrackRule{
{Action: ConntrackNoTrack, Source: "net", User: "nobody"},
},
wantErr: "user match only valid for output chain",
},
{
name: "user with output chain is valid",
rules: []ConntrackRule{
{Action: ConntrackNoTrack, Source: "net", User: "nobody", Chain: ConntrackOutput},
},
},
{
name: "source or dest required for non-helper",
rules: []ConntrackRule{
{Action: ConntrackDrop},
},
wantErr: "source or dest required",
},
{
name: "helper without source/dest is valid",
rules: []ConntrackRule{
{Action: ConntrackHelper, Helper: "ftp", Proto: "tcp", Chain: ConntrackBoth},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Conntrack = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 2. Blrules ---
func TestValidateBlrules(t *testing.T) {
tests := []struct {
name string
rules []BlruleRule
wantErr string
}{
{
name: "valid rule",
rules: []BlruleRule{
{Action: BlruleAccept, Source: "net", Dest: "loc"},
},
},
{
name: "unknown action",
rules: []BlruleRule{
{Action: "bogus", Source: "net", Dest: "loc"},
},
wantErr: "unknown action",
},
{
name: "source zone not defined",
rules: []BlruleRule{
{Action: BlruleDrop, Source: "nosuchzone", Dest: "loc"},
},
wantErr: `source zone "nosuchzone" not defined`,
},
{
name: "source required",
rules: []BlruleRule{
{Action: BlruleDrop, Dest: "loc"},
},
wantErr: "source required",
},
{
name: "dest required",
rules: []BlruleRule{
{Action: BlruleDrop, Source: "net"},
},
wantErr: "dest required",
},
{
name: "source all is valid",
rules: []BlruleRule{
{Action: BlruleDrop, Source: "all", Dest: "all"},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Blrules = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 3. Tunnels ---
func TestValidateTunnels(t *testing.T) {
tests := []struct {
name string
tunnels []Tunnel
wantErr string
}{
{
name: "valid tunnel",
tunnels: []Tunnel{
{Type: "ipsec", Zone: "net", Gateways: []string{"1.2.3.4"}},
},
},
{
name: "zone not defined",
tunnels: []Tunnel{
{Type: "gre", Zone: "nosuchzone", Gateways: []string{"1.2.3.4"}},
},
wantErr: `zone "nosuchzone" not defined`,
},
{
name: "missing gateways",
tunnels: []Tunnel{
{Type: "ipsec", Zone: "net"},
},
wantErr: "at least one gateway required",
},
{
name: "unknown tunnel type",
tunnels: []Tunnel{
{Type: "bogus", Zone: "net", Gateways: []string{"1.2.3.4"}},
},
wantErr: "unknown tunnel type",
},
{
name: "tunnel type with suffix",
tunnels: []Tunnel{
{Type: "ipsec:ah", Zone: "net", Gateways: []string{"1.2.3.4"}},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Tunnels = tt.tunnels
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 4. Rtrules ---
func TestValidateRoutingRules(t *testing.T) {
withProvider := func() Config {
cfg := baseConfig()
cfg.Providers = []Provider{
{Name: "isp1", Number: 1, Interface: "eth0"},
}
return cfg
}
tests := []struct {
name string
setup func() Config
rules []RoutingRule
wantErr string
}{
{
name: "valid routing rule",
setup: withProvider,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Provider: "isp1", Priority: 1000},
},
},
{
name: "missing providers",
setup: baseConfig,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Provider: "isp1", Priority: 1000},
},
wantErr: "rtrules require providers",
},
{
name: "provider not defined",
setup: withProvider,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Provider: "nosuch", Priority: 1000},
},
wantErr: `provider "nosuch" not defined`,
},
{
name: "priority below range",
setup: withProvider,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Provider: "isp1", Priority: 999},
},
wantErr: "priority must be 1000-26999",
},
{
name: "priority above range",
setup: withProvider,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Provider: "isp1", Priority: 27000},
},
wantErr: "priority must be 1000-26999",
},
{
name: "source or dest required",
setup: withProvider,
rules: []RoutingRule{
{Provider: "isp1", Priority: 1000},
},
wantErr: "source or dest required",
},
{
name: "provider required",
setup: withProvider,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Priority: 1000},
},
wantErr: "provider required",
},
{
name: "main provider always valid",
setup: withProvider,
rules: []RoutingRule{
{Source: "10.0.0.0/8", Provider: "main", Priority: 1000},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := tt.setup()
cfg.RoutingRules = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 5. StoppedRules ---
func TestValidateStoppedRules(t *testing.T) {
tests := []struct {
name string
rules []StoppedRule
wantErr string
}{
{
name: "valid rule",
rules: []StoppedRule{
{Action: StoppedAccept, Source: "eth0"},
},
},
{
name: "unknown action",
rules: []StoppedRule{
{Action: "bogus", Source: "eth0"},
},
wantErr: "unknown action",
},
{
name: "missing source and dest",
rules: []StoppedRule{
{Action: StoppedAccept},
},
wantErr: "source or dest required",
},
{
name: "accept with source and dest is valid",
rules: []StoppedRule{
{Action: StoppedAccept, Source: "eth0", Dest: "eth1"},
},
},
{
name: "notrack with dest $FW is valid",
rules: []StoppedRule{
{Action: StoppedNoTrack, Source: "eth0", Dest: "$FW"},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.StoppedRules = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 6. Vars ---
func TestValidateVars(t *testing.T) {
tests := []struct {
name string
vars map[string]string
wantErr string
}{
{
name: "valid vars",
vars: map[string]string{"NET_IF": "eth0", "PORT": "8080"},
},
{
name: "empty var name",
vars: map[string]string{"": "value"},
wantErr: "empty variable name",
},
{
name: "invalid chars - space",
vars: map[string]string{"bad name": "value"},
wantErr: "invalid variable name",
},
{
name: "invalid chars - dollar",
vars: map[string]string{"$var": "value"},
wantErr: "invalid variable name",
},
{
name: "invalid chars - braces",
vars: map[string]string{"{var}": "value"},
wantErr: "invalid variable name",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Vars = tt.vars
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 7. NAT (static) ---
func TestValidateStaticNAT(t *testing.T) {
tests := []struct {
name string
nat []StaticNAT
wantErr string
}{
{
name: "valid nat",
nat: []StaticNAT{
{External: "1.2.3.4", Interface: "eth0", Internal: "10.0.0.1"},
},
},
{
name: "external required",
nat: []StaticNAT{
{Interface: "eth0", Internal: "10.0.0.1"},
},
wantErr: "external address required",
},
{
name: "external must be IP",
nat: []StaticNAT{
{External: "example.com", Interface: "eth0", Internal: "10.0.0.1"},
},
wantErr: "external must be an IP address",
},
{
name: "interface required",
nat: []StaticNAT{
{External: "1.2.3.4", Internal: "10.0.0.1"},
},
wantErr: "interface required",
},
{
name: "internal required",
nat: []StaticNAT{
{External: "1.2.3.4", Interface: "eth0"},
},
wantErr: "internal address required",
},
{
name: "internal must be IP",
nat: []StaticNAT{
{External: "1.2.3.4", Interface: "eth0", Internal: "server.local"},
},
wantErr: "internal must be an IP address",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.StaticNAT = tt.nat
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 8. Netmap ---
func TestValidateNetmap(t *testing.T) {
tests := []struct {
name string
netmap []Netmap
wantErr string
}{
{
name: "valid netmap",
netmap: []Netmap{
{Type: NetmapDNAT, Net1: "192.168.1.0/24", Interface: "eth0", Net2: "10.0.0.0/24"},
},
},
{
name: "invalid CIDR net1",
netmap: []Netmap{
{Type: NetmapDNAT, Net1: "notacidr", Interface: "eth0", Net2: "10.0.0.0/24"},
},
wantErr: "net1 must be CIDR format",
},
{
name: "invalid CIDR net2",
netmap: []Netmap{
{Type: NetmapSNAT, Net1: "192.168.1.0/24", Interface: "eth0", Net2: "notacidr"},
},
wantErr: "net2 must be CIDR format",
},
{
name: "missing interface",
netmap: []Netmap{
{Type: NetmapDNAT, Net1: "192.168.1.0/24", Net2: "10.0.0.0/24"},
},
wantErr: "interface required",
},
{
name: "interface not defined",
netmap: []Netmap{
{Type: NetmapDNAT, Net1: "192.168.1.0/24", Interface: "eth99", Net2: "10.0.0.0/24"},
},
wantErr: `interface "eth99" not defined`,
},
{
name: "invalid type",
netmap: []Netmap{
{Type: "bogus", Net1: "192.168.1.0/24", Interface: "eth0", Net2: "10.0.0.0/24"},
},
wantErr: "type must be dnat or snat",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Netmap = tt.netmap
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 9. Providers ---
func TestValidateProviders(t *testing.T) {
tests := []struct {
name string
providers []Provider
wantErr string
}{
{
name: "valid provider",
providers: []Provider{
{Name: "isp1", Number: 1, Interface: "eth0"},
},
},
{
name: "duplicate name",
providers: []Provider{
{Name: "isp1", Number: 1, Interface: "eth0"},
{Name: "isp1", Number: 2, Interface: "eth1"},
},
wantErr: `duplicate name "isp1"`,
},
{
name: "reserved name local",
providers: []Provider{
{Name: "local", Number: 1, Interface: "eth0"},
},
wantErr: `"local" is a reserved name`,
},
{
name: "reserved name main",
providers: []Provider{
{Name: "main", Number: 1, Interface: "eth0"},
},
wantErr: `"main" is a reserved name`,
},
{
name: "number below range",
providers: []Provider{
{Name: "isp1", Number: 0, Interface: "eth0"},
},
wantErr: "number must be between 1 and 252",
},
{
name: "number above range",
providers: []Provider{
{Name: "isp1", Number: 253, Interface: "eth0"},
},
wantErr: "number must be between 1 and 252",
},
{
name: "duplicate number",
providers: []Provider{
{Name: "isp1", Number: 1, Interface: "eth0"},
{Name: "isp2", Number: 1, Interface: "eth1"},
},
wantErr: "number 1 already used",
},
{
name: "interface required",
providers: []Provider{
{Name: "isp1", Number: 1},
},
wantErr: "interface required",
},
{
name: "name required",
providers: []Provider{
{Number: 1, Interface: "eth0"},
},
wantErr: "name required",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Providers = tt.providers
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 10. Accounting ---
func TestValidateAccounting(t *testing.T) {
tests := []struct {
name string
rules []AccountingRule
wantErr string
}{
{
name: "valid rule",
rules: []AccountingRule{
{Action: AccountingCount, Section: AccountingSectionForward, Source: "net"},
},
},
{
name: "source or dest required",
rules: []AccountingRule{
{Action: AccountingCount, Section: AccountingSectionForward},
},
wantErr: "source or dest required",
},
{
name: "unknown action",
rules: []AccountingRule{
{Action: "bogus", Section: AccountingSectionForward, Source: "net"},
},
wantErr: "unknown action",
},
{
name: "unknown section",
rules: []AccountingRule{
{Action: AccountingCount, Section: "bogus", Source: "net"},
},
wantErr: "unknown section",
},
{
name: "dest only is valid",
rules: []AccountingRule{
{Action: AccountingDone, Section: AccountingSectionInput, Dest: "loc"},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Accounting = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 11. Mangle ---
func TestValidateMangle(t *testing.T) {
tests := []struct {
name string
rules []MangleRule
wantErr string
}{
{
name: "valid rule",
rules: []MangleRule{
{Action: MangleMark, Chain: ManglePrerouting, MarkValue: "0x1"},
},
},
{
name: "mark_value required for mark action",
rules: []MangleRule{
{Action: MangleMark, Chain: ManglePrerouting},
},
wantErr: "mark_value required for mark action",
},
{
name: "mark_value required for connmark action",
rules: []MangleRule{
{Action: MangleConnMark, Chain: ManglePrerouting},
},
wantErr: "mark_value required for connmark action",
},
{
name: "mark_value required for classify action",
rules: []MangleRule{
{Action: MangleClassify, Chain: MangleForward},
},
wantErr: "mark_value required for classify action",
},
{
name: "unknown action",
rules: []MangleRule{
{Action: "bogus", Chain: ManglePrerouting},
},
wantErr: "unknown action",
},
{
name: "unknown chain",
rules: []MangleRule{
{Action: MangleDrop, Chain: "bogus"},
},
wantErr: "unknown chain",
},
{
name: "log action without mark_value is valid",
rules: []MangleRule{
{Action: MangleLog, Chain: MangleInput},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Mangle = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 12. Maclist ---
func TestValidateMaclist(t *testing.T) {
tests := []struct {
name string
entries []MaclistEntry
wantErr string
}{
{
name: "valid entry with mac",
entries: []MaclistEntry{
{Action: MaclistAccept, Interface: "eth0", MAC: "00:11:22:33:44:55"},
},
},
{
name: "valid entry with addresses",
entries: []MaclistEntry{
{Action: MaclistDrop, Interface: "eth0", Addresses: []string{"10.0.0.1"}},
},
},
{
name: "mac or addresses required",
entries: []MaclistEntry{
{Action: MaclistAccept, Interface: "eth0"},
},
wantErr: "mac or addresses required",
},
{
name: "interface required",
entries: []MaclistEntry{
{Action: MaclistAccept, MAC: "00:11:22:33:44:55"},
},
wantErr: "interface required",
},
{
name: "unknown action",
entries: []MaclistEntry{
{Action: "bogus", Interface: "eth0", MAC: "00:11:22:33:44:55"},
},
wantErr: "unknown action",
},
{
name: "interface not defined",
entries: []MaclistEntry{
{Action: MaclistAccept, Interface: "eth99", MAC: "00:11:22:33:44:55"},
},
wantErr: `interface "eth99" not defined`,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Maclist = tt.entries
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 13. TC ---
func TestValidateTCDevices(t *testing.T) {
tests := []struct {
name string
devices []TCDevice
wantErr string
}{
{
name: "valid device",
devices: []TCDevice{
{Interface: "eth0", OutBandwidth: "10mbit"},
},
},
{
name: "interface required",
devices: []TCDevice{
{OutBandwidth: "10mbit"},
},
wantErr: "interface required",
},
{
name: "out_bandwidth required",
devices: []TCDevice{
{Interface: "eth0"},
},
wantErr: "out_bandwidth required",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.TCDevices = tt.devices
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
func TestValidateTCClasses(t *testing.T) {
tests := []struct {
name string
classes []TCClass
wantErr string
}{
{
name: "valid class",
classes: []TCClass{
{Interface: "eth0:1", Rate: "1mbit"},
},
},
{
name: "interface required",
classes: []TCClass{
{Rate: "1mbit"},
},
wantErr: "interface required",
},
{
name: "rate required",
classes: []TCClass{
{Interface: "eth0:1"},
},
wantErr: "rate required",
},
{
name: "mark out of range",
classes: []TCClass{
{Interface: "eth0:1", Rate: "1mbit", Mark: 256},
},
wantErr: "mark must be 1-255",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.TCClasses = tt.classes
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
func TestValidateTCFilters(t *testing.T) {
tests := []struct {
name string
filters []TCFilter
wantErr string
}{
{
name: "valid filter",
filters: []TCFilter{
{Class: "eth0:1", Source: "10.0.0.0/8"},
},
},
{
name: "class required",
filters: []TCFilter{
{Source: "10.0.0.0/8"},
},
wantErr: "class required",
},
{
name: "source or dest required",
filters: []TCFilter{
{Class: "eth0:1"},
},
wantErr: "source or dest required",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.TCFilters = tt.filters
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
func TestValidateTCInterfaces(t *testing.T) {
tests := []struct {
name string
interfaces []TCInterface
wantErr string
}{
{
name: "valid interface",
interfaces: []TCInterface{
{Interface: "eth0"},
},
},
{
name: "interface required",
interfaces: []TCInterface{
{Type: "external"},
},
wantErr: "interface required",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.TCInterfaces = tt.interfaces
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
func TestValidateTCPriority(t *testing.T) {
tests := []struct {
name string
priorities []TCPriority
wantErr string
}{
{
name: "valid band 1",
priorities: []TCPriority{
{Band: 1},
},
},
{
name: "valid band 3",
priorities: []TCPriority{
{Band: 3},
},
},
{
name: "band below range",
priorities: []TCPriority{
{Band: 0},
},
wantErr: "band must be 1, 2, or 3",
},
{
name: "band above range",
priorities: []TCPriority{
{Band: 4},
},
wantErr: "band must be 1, 2, or 3",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.TCPriorities = tt.priorities
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 14. ProxyARP ---
func TestValidateProxyARP(t *testing.T) {
tests := []struct {
name string
entries []ProxyARP
wantErr string
}{
{
name: "valid entry",
entries: []ProxyARP{
{Address: "1.2.3.4", Interface: "eth1", External: "eth0"},
},
},
{
name: "address required",
entries: []ProxyARP{
{Interface: "eth1", External: "eth0"},
},
wantErr: "address required",
},
{
name: "external required",
entries: []ProxyARP{
{Address: "1.2.3.4", Interface: "eth1"},
},
wantErr: "external required",
},
{
name: "interface required unless haveroute",
entries: []ProxyARP{
{Address: "1.2.3.4", External: "eth0"},
},
wantErr: "interface required unless haveroute",
},
{
name: "haveroute skips interface requirement",
entries: []ProxyARP{
{Address: "1.2.3.4", External: "eth0", HaveRoute: true},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.ProxyARP = tt.entries
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 15. Routes ---
func TestValidateRoutes(t *testing.T) {
tests := []struct {
name string
routes []StaticRoute
wantErr string
}{
{
name: "valid route",
routes: []StaticRoute{
{Provider: "main", Dest: "10.0.0.0/8", Gateway: "192.168.1.1"},
},
},
{
name: "provider required",
routes: []StaticRoute{
{Dest: "10.0.0.0/8", Gateway: "192.168.1.1"},
},
wantErr: "provider required",
},
{
name: "dest required",
routes: []StaticRoute{
{Provider: "main", Gateway: "192.168.1.1"},
},
wantErr: "dest required",
},
{
name: "device not allowed with blackhole",
routes: []StaticRoute{
{Provider: "main", Dest: "10.0.0.0/8", Gateway: "blackhole", Device: "eth0"},
},
wantErr: "device not allowed with blackhole gateway",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Routes = tt.routes
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 16. ArpRules ---
func TestValidateArpRules(t *testing.T) {
tests := []struct {
name string
rules []ArpRule
wantErr string
}{
{
name: "valid rule",
rules: []ArpRule{
{Action: ArpAccept, Source: "net"},
},
},
{
name: "action_address required for snat",
rules: []ArpRule{
{Action: ArpSNAT, Source: "net"},
},
wantErr: "action_address required for snat action",
},
{
name: "action_address required for dnat",
rules: []ArpRule{
{Action: ArpDNAT, Source: "net"},
},
wantErr: "action_address required for dnat action",
},
{
name: "action_mac required for smat",
rules: []ArpRule{
{Action: ArpSMAT, Source: "net"},
},
wantErr: "action_mac required for smat action",
},
{
name: "unknown action",
rules: []ArpRule{
{Action: "bogus", Source: "net"},
},
wantErr: "unknown action",
},
{
name: "source or dest required",
rules: []ArpRule{
{Action: ArpDrop},
},
wantErr: "source or dest required",
},
{
name: "valid snat with action_address",
rules: []ArpRule{
{Action: ArpSNAT, Source: "net", ActionAddress: "1.2.3.4"},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.ArpRules = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 17. Secmarks ---
func TestValidateSecmarks(t *testing.T) {
tests := []struct {
name string
rules []SecmarkRule
wantErr string
}{
{
name: "valid secmark",
rules: []SecmarkRule{
{Secmark: "system_u:object_r:httpd_t:s0", Chain: "P"},
},
},
{
name: "secmark required",
rules: []SecmarkRule{
{Chain: "P"},
},
wantErr: "secmark required",
},
{
name: "chain required",
rules: []SecmarkRule{
{Secmark: "system_u:object_r:httpd_t:s0"},
},
wantErr: "chain required",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Secmarks = tt.rules
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 18. Zone nesting ---
func TestResolveNesting(t *testing.T) {
t.Run("simple parent-child hierarchy", func(t *testing.T) {
cfg := Config{
Zones: map[string]Zone{
"fw": {Type: ZoneFirewall},
"net": {Type: ZoneIP},
"dmz": {Type: ZoneIP, Parents: []string{"net"}},
},
}
order, err := cfg.ResolveNesting()
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
// dmz (child) must appear before net (parent) or at least both must appear.
// The key constraint: all zones should be present.
if len(order) != 3 {
t.Fatalf("expected 3 zones in order, got %d: %v", len(order), order)
}
})
t.Run("circular detection", func(t *testing.T) {
// The algorithm requires all siblings (children of the same parent)
// to be resolved before any of them can proceed. Two children of
// the same non-firewall parent create a deadlock that is reported
// as a circular nesting error.
cfg := Config{
Zones: map[string]Zone{
"fw": {Type: ZoneFirewall},
"parent": {Type: ZoneIP},
"child1": {Type: ZoneIP, Parents: []string{"parent"}},
"child2": {Type: ZoneIP, Parents: []string{"parent"}},
},
}
_, err := cfg.ResolveNesting()
if err == nil {
t.Fatal("expected circular nesting error")
}
if !strings.Contains(err.Error(), "circular") {
t.Fatalf("expected circular error, got: %v", err)
}
})
}
func TestChildZones(t *testing.T) {
cfg := Config{
Zones: map[string]Zone{
"fw": {Type: ZoneFirewall},
"net": {Type: ZoneIP},
"dmz": {Type: ZoneIP, Parents: []string{"net"}},
"vpn": {Type: ZoneIP, Parents: []string{"net"}},
"loc": {Type: ZoneIP},
},
}
children := cfg.ChildZones("net")
if len(children) != 2 {
t.Fatalf("expected 2 children of net, got %d: %v", len(children), children)
}
// Check both dmz and vpn are present.
found := map[string]bool{}
for _, c := range children {
found[c] = true
}
if !found["dmz"] || !found["vpn"] {
t.Fatalf("expected dmz and vpn as children, got: %v", children)
}
// loc has no children.
locChildren := cfg.ChildZones("loc")
if len(locChildren) != 0 {
t.Fatalf("expected 0 children of loc, got %d", len(locChildren))
}
}
func TestIsSubZone(t *testing.T) {
cfg := Config{
Zones: map[string]Zone{
"fw": {Type: ZoneFirewall},
"net": {Type: ZoneIP},
"dmz": {Type: ZoneIP, Parents: []string{"net"}},
"web": {Type: ZoneIP, Parents: []string{"dmz"}},
},
}
tests := []struct {
child, parent string
want bool
}{
{"dmz", "net", true},
{"web", "dmz", true},
{"web", "net", true}, // transitive
{"net", "dmz", false}, // reverse
{"net", "net", false}, // self
{"nosuch", "net", false}, // non-existent
}
for _, tt := range tests {
t.Run(tt.child+"->"+tt.parent, func(t *testing.T) {
got := cfg.IsSubZone(tt.child, tt.parent)
if got != tt.want {
t.Fatalf("IsSubZone(%q, %q) = %v, want %v", tt.child, tt.parent, got, tt.want)
}
})
}
}
// --- 19. Names ---
func TestValidateName(t *testing.T) {
tests := []struct {
name string
input string
wantErr string
}{
{"valid simple", "net", ""},
{"valid with underscore", "my_zone", ""},
{"valid with digits", "zone1", ""},
{"empty name", "", "empty"},
{"starts with digit", "1zone", "must start with a letter"},
{"contains dash", "my-zone", "invalid character"},
{"contains space", "my zone", "invalid character"},
{"contains dot", "my.zone", "invalid character"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := ValidateName(tt.input, "test")
checkErr(t, err, tt.wantErr)
})
}
}
// --- 20. Params ---
func TestSubstituteVars(t *testing.T) {
tests := []struct {
name string
input string
vars map[string]string
want string
}{
{
name: "braced substitution",
input: "${NET_IF}",
vars: map[string]string{"NET_IF": "eth0"},
want: "eth0",
},
{
name: "unbraced substitution",
input: "$NET_IF",
vars: map[string]string{"NET_IF": "eth0"},
want: "eth0",
},
{
name: "no vars",
input: "no substitution",
vars: nil,
want: "no substitution",
},
{
name: "no dollar sign",
input: "plain text",
vars: map[string]string{"foo": "bar"},
want: "plain text",
},
{
name: "multiple vars",
input: "${A}:${B}",
vars: map[string]string{"A": "1", "B": "2"},
want: "1:2",
},
{
name: "undefined var stays",
input: "${UNDEF}",
vars: map[string]string{"OTHER": "val"},
want: "${UNDEF}",
},
{
name: "empty vars map",
input: "$foo",
vars: map[string]string{},
want: "$foo",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := SubstituteVars(tt.input, tt.vars)
if got != tt.want {
t.Fatalf("SubstituteVars(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
// --- 22. ProxyNDP ---
func TestValidateProxyNDP(t *testing.T) {
tests := []struct {
name string
entries []ProxyNDP
wantErr string
}{
{
name: "valid entry",
entries: []ProxyNDP{
{Address: "fd10::100", Interface: "eth0", External: "eth1"},
},
},
{
name: "address required",
entries: []ProxyNDP{
{Interface: "eth0", External: "eth1"},
},
wantErr: "address required",
},
{
name: "external required",
entries: []ProxyNDP{
{Address: "fd10::100", Interface: "eth0"},
},
wantErr: "external required",
},
{
name: "interface required unless haveroute",
entries: []ProxyNDP{
{Address: "fd10::100", External: "eth1"},
},
wantErr: "interface required unless haveroute",
},
{
name: "haveroute skips interface requirement",
entries: []ProxyNDP{
{Address: "fd10::100", External: "eth1", HaveRoute: true},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.ProxyNDP = tt.entries
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}
// --- 23. Settings ---
func TestValidateSettings(t *testing.T) {
tests := []struct {
name string
family AddressFamily
wantErr string
}{
{"inet is valid", FamilyINET, ""},
{"ip is valid", FamilyIP, ""},
{"ip6 is valid", FamilyIP6, ""},
{"empty is invalid", "", "unknown address_family"},
{"bogus is invalid", "bogus", "unknown address_family"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cfg := baseConfig()
cfg.Settings.AddressFamily = tt.family
err := cfg.Validate()
checkErr(t, err, tt.wantErr)
})
}
}