Merge benvin/multiport-multiproto; reject limits on zone and address lists
This commit is contained in:
@@ -268,6 +268,16 @@ func (c *Compiler) compileRules(state *FirewallState) error {
|
||||
sport = rule.SPort
|
||||
}
|
||||
|
||||
if rule.RateLimit != "" || rule.ConnLimit != "" {
|
||||
matches, err := l4Matches(proto, dports, sport)
|
||||
if err != nil {
|
||||
return fmt.Errorf("rule[%d]: %w", i, err)
|
||||
}
|
||||
if len(matches)*specCount(rule.Source, rule.Dest, rule.Action) > 1 {
|
||||
return fmt.Errorf("rule[%d]: ratelimit/connlimit cannot be combined with proto, port, zone or address lists (each expanded rule would get its own limiter)", i)
|
||||
}
|
||||
}
|
||||
|
||||
if err := c.compileOneRule(state, tag, rule.Source, rule.Dest,
|
||||
proto, dports, sport,
|
||||
rule.Action, rule.Log, rule.Dest, fwZone, rule.Section); err != nil {
|
||||
@@ -368,6 +378,20 @@ func (c *Compiler) compileOneRule(state *FirewallState, tag, srcSpec, dstSpec, p
|
||||
return nil
|
||||
}
|
||||
|
||||
// specCount is how many zone/address combinations compileOneRule expands src and dst into.
|
||||
func specCount(srcSpec, dstSpec string, action config.RuleAction) int {
|
||||
count := func(spec string) (n int) {
|
||||
for _, z := range zoneSpecs(spec) {
|
||||
n += len(splitAddrs(z.Addr))
|
||||
}
|
||||
return n
|
||||
}
|
||||
if action == config.RuleDNAT || action == config.RuleRedirect {
|
||||
return count(srcSpec)
|
||||
}
|
||||
return count(srcSpec) * count(dstSpec)
|
||||
}
|
||||
|
||||
// zoneSpecs expands a comma zone list; "all"/"any" forms keep their own comma (exclusion) syntax.
|
||||
func zoneSpecs(spec string) []config.ZoneSpec {
|
||||
if strings.HasPrefix(spec, "all") || strings.HasPrefix(spec, "any") {
|
||||
@@ -1006,10 +1030,13 @@ func l4Matches(proto string, dports, sports config.PortSpec) ([]l4Match, error)
|
||||
var out []l4Match
|
||||
for _, p := range protos {
|
||||
p = strings.TrimSpace(p)
|
||||
if proto != "" && p == "" {
|
||||
return nil, fmt.Errorf("empty element in proto list %q", proto)
|
||||
}
|
||||
isICMP := strings.EqualFold(p, "icmp") || strings.EqualFold(p, "icmpv6") || strings.EqualFold(p, "ipv6-icmp")
|
||||
parseD := parsePortOrRange
|
||||
if isICMP {
|
||||
parseD = func(s string) ([]expr.Any, error) { return matchICMPType(s), nil }
|
||||
parseD = matchICMPType
|
||||
}
|
||||
dalts, err := portAlternatives(dports, parseD)
|
||||
if err != nil {
|
||||
@@ -1038,7 +1065,7 @@ func portAlternatives(ports config.PortSpec, parse func(string) ([]expr.Any, err
|
||||
for _, item := range ports {
|
||||
for _, s := range strings.Split(item, ",") {
|
||||
if s = strings.TrimSpace(s); s == "" {
|
||||
continue
|
||||
return nil, fmt.Errorf("empty element in port list %q", item)
|
||||
}
|
||||
e, err := parse(s)
|
||||
if err != nil {
|
||||
@@ -1216,33 +1243,33 @@ var icmpTypeNames = map[string]byte{
|
||||
"address-mask-reply": 18,
|
||||
}
|
||||
|
||||
func matchICMPType(spec string) []expr.Any {
|
||||
func matchICMPType(spec string) ([]expr.Any, error) {
|
||||
if strings.Contains(spec, "/") {
|
||||
parts := strings.SplitN(spec, "/", 2)
|
||||
typeVal, ok := resolveICMPType(parts[0])
|
||||
if !ok {
|
||||
return nil
|
||||
return nil, fmt.Errorf("invalid icmp type %q", parts[0])
|
||||
}
|
||||
code, err := strconv.ParseUint(parts[1], 10, 8)
|
||||
if err != nil {
|
||||
return nil
|
||||
return nil, fmt.Errorf("invalid icmp code %q: %w", parts[1], err)
|
||||
}
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}},
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 1, Len: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{byte(code)}},
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
typeVal, ok := resolveICMPType(spec)
|
||||
if !ok {
|
||||
return nil
|
||||
return nil, fmt.Errorf("invalid icmp type %q", spec)
|
||||
}
|
||||
return []expr.Any{
|
||||
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{typeVal}},
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
func resolveICMPType(s string) (byte, bool) {
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/google/nftables/expr"
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"git.unkin.net/unkin/tomswall/internal/config"
|
||||
)
|
||||
@@ -1269,7 +1270,10 @@ func TestMatchICMPType(t *testing.T) {
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
exprs := matchICMPType(tt.input)
|
||||
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)
|
||||
}
|
||||
@@ -1729,37 +1733,102 @@ func TestCompile_CommaZoneLists(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompile_CommaZoneListExtras(t *testing.T) {
|
||||
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}, "lan": {Type: config.ZoneIP}},
|
||||
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}},
|
||||
Rules: []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"}},
|
||||
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),
|
||||
}
|
||||
state, err := NewCompiler(cfg).Compile()
|
||||
mod(cfg)
|
||||
return cfg
|
||||
}
|
||||
|
||||
func taggedRules(state *FirewallState, chain, tag string) []ManagedRule {
|
||||
var out []ManagedRule
|
||||
for _, r := range state.Rules[chain] {
|
||||
if r.Tag == tag {
|
||||
out = append(out, r)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func TestCompile_ListExpansionCounts(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mod func(*config.Config)
|
||||
chain string
|
||||
tag string
|
||||
want int
|
||||
}{
|
||||
{"proto x dport cross product", func(c *config.Config) {
|
||||
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80", "443"}}}
|
||||
}, "input", "rule:0", 4},
|
||||
{"sport list", func(c *config.Config) {
|
||||
c.Rules = []config.Rule{{Action: config.RuleAccept, Source: "net", Dest: "fw", Proto: "tcp", SPort: config.PortSpec{"1024,2048"}}}
|
||||
}, "input", "rule:0", 2},
|
||||
{"snat proto x dport", func(c *config.Config) {
|
||||
c.SNAT = []config.SNATRule{{Action: config.SNATMasquerade, Dest: "eth0", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}}
|
||||
}, "postrouting", "snat:0", 4},
|
||||
{"conntrack dport list", func(c *config.Config) {
|
||||
c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"53", "123"}}}
|
||||
}, "prerouting", "conntrack:0:prerouting", 2},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
state, err := NewCompiler(listCfg(tt.mod)).Compile()
|
||||
if err != nil {
|
||||
t.Fatalf("Compile() error: %v", err)
|
||||
}
|
||||
if got := len(taggedRules(state, tt.chain, tt.tag)); got != tt.want {
|
||||
t.Errorf("%s rules in %s = %d, want %d", tt.tag, tt.chain, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompile_CommaZoneListLimitErrors(t *testing.T) {
|
||||
for _, r := range []config.Rule{
|
||||
{Action: config.RuleAccept, Source: "net", Dest: "fw,lan", RateLimit: "10/sec:5"},
|
||||
{Action: config.RuleAccept, Source: "net,lan", Dest: "fw", ConnLimit: "10"},
|
||||
{Action: config.RuleAccept, Source: "net", Dest: "fw:192.0.2.1,198.51.100.1", RateLimit: "10/sec"},
|
||||
} {
|
||||
t.Run(r.Source+">"+r.Dest, func(t *testing.T) {
|
||||
cfg := &config.Config{
|
||||
Settings: config.Settings{TableName: "test", AddressFamily: config.FamilyINET},
|
||||
Zones: map[string]config.Zone{"fw": {Type: config.ZoneFirewall}, "net": {Type: config.ZoneIP}, "lan": {Type: config.ZoneIP}},
|
||||
Interfaces: []config.Interface{{Zone: "net", Interface: "eth0"}, {Zone: "lan", Interface: "eth1"}},
|
||||
Rules: []config.Rule{r},
|
||||
PortGroups: make(map[string]config.PortGroup),
|
||||
}
|
||||
if _, err := NewCompiler(cfg).Compile(); err == nil {
|
||||
t.Fatal("Compile() succeeded, want error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompile_RejectPerProto(t *testing.T) {
|
||||
state, err := NewCompiler(listCfg(func(c *config.Config) {
|
||||
c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"53"}}}
|
||||
})).Compile()
|
||||
if err != nil {
|
||||
t.Fatalf("Compile() error: %v", err)
|
||||
}
|
||||
for _, chain := range []string{"input", "forward"} {
|
||||
found := false
|
||||
for _, r := range state.Rules[chain] {
|
||||
if r.Tag != "rule:0" {
|
||||
continue
|
||||
}
|
||||
found = true
|
||||
hasLimit := false
|
||||
for _, e := range r.Exprs {
|
||||
if _, ok := e.(*expr.Limit); ok {
|
||||
hasLimit = true
|
||||
}
|
||||
}
|
||||
if !hasLimit {
|
||||
t.Errorf("%s rule missing Limit expression", chain)
|
||||
}
|
||||
rules := taggedRules(state, "input", "rule:0")
|
||||
if len(rules) != 2 {
|
||||
t.Fatalf("got %d rules, want 2", len(rules))
|
||||
}
|
||||
want := []uint32{unix.NFT_REJECT_TCP_RST, unix.NFT_REJECT_ICMPX_UNREACH}
|
||||
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 !found {
|
||||
t.Errorf("rule:0 not found in %s chain", chain)
|
||||
if rej.Type != want[i] {
|
||||
t.Errorf("rule %d: reject type %d, want %d", i, rej.Type, want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1778,3 +1847,44 @@ func TestNegatedAddressList(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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"}}},
|
||||
{"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_ListExpansionDiffStable(t *testing.T) {
|
||||
cfg := listCfg(func(c *config.Config) {
|
||||
c.Rules = []config.Rule{{Action: config.RuleReject, Source: "net", Dest: "fw", Proto: "tcp,udp", DPort: config.PortSpec{"80,443"}}}
|
||||
})
|
||||
state, err := NewCompiler(cfg).Compile()
|
||||
if err != nil {
|
||||
t.Fatalf("Compile() error: %v", err)
|
||||
}
|
||||
if n := len(taggedRules(state, "input", "rule:0")); n != 4 {
|
||||
t.Fatalf("got %d rule:0 rules, want 4", n)
|
||||
}
|
||||
if cs := computeDiff(state, state); !cs.Empty() {
|
||||
t.Errorf("diff not empty:\n%s", cs.Summary())
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user