From be391ed385688be3bf7023c7434955c8a48daceb Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sun, 4 Oct 2026 15:56:23 +1100 Subject: [PATCH] Keep rule match extras on the limited log rule --- internal/nftables/compiler.go | 47 +++++++++------- internal/nftables/compiler_test.go | 87 ++++++++++++++++++++++++++++++ 2 files changed, 114 insertions(+), 20 deletions(-) diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 1440a03..47244d4 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -72,7 +72,8 @@ func (c *Compiler) Compile() (*FirewallState, error) { // limitLogs puts a limit in front of every log expression. A limit stops the // whole rule, so like shorewall's separate LOG rule, a log followed by an action -// splits into a limited log-only rule and the same rule without the log. +// splits into a limited log-only rule and the same rule without the log. A +// LOG-action rule (nothing but rate limits after the log) keeps only the log rule. func limitLogs(state *FirewallState, spec string) { if spec == "" { return @@ -88,7 +89,7 @@ func limitLogs(state *FirewallState, spec string) { logRule := r logRule.Exprs = slices.Concat(r.Exprs[:i], parseRateLimit(spec), r.Exprs[i:i+1]) out = append(out, logRule) - if i < len(r.Exprs)-1 { + if slices.ContainsFunc(r.Exprs[i+1:], func(e expr.Any) bool { _, ok := e.(*expr.Limit); return !ok }) { r.Exprs = slices.Concat(r.Exprs[:i], r.Exprs[i+1:]) out = append(out, r) } @@ -442,24 +443,24 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul break } - var extra []expr.Any + var match, extra []expr.Any var replaceVerdict []expr.Any if rule.User != "" { - extra = append(extra, matchUID(rule.User)...) + match = append(match, matchUID(rule.User)...) } if rule.Mark != "" { - extra = append(extra, matchMark(rule.Mark)...) + match = append(match, matchMark(rule.Mark)...) + } + if rule.ConnLimit != "" { + match = append(match, matchConnLimit(rule.ConnLimit)...) + } + if rule.Time != nil { + match = append(match, matchTime(rule.Time)...) } if rule.RateLimit != "" { extra = append(extra, parseRateLimit(rule.RateLimit)...) } - if rule.ConnLimit != "" { - extra = append(extra, matchConnLimit(rule.ConnLimit)...) - } - if rule.Time != nil { - extra = append(extra, matchTime(rule.Time)...) - } if rule.SetMark != "" { extra = append(extra, setMarkExprs(rule.SetMark)...) } @@ -467,21 +468,27 @@ func (c *Compiler) applyChainExtras(state *FirewallState, chain, tag string, rul replaceVerdict = []expr.Any{&expr.Queue{Num: uint16(rule.NFQueue), Total: 1}} } - if len(extra) > 0 || len(replaceVerdict) > 0 { - existingExprs := rules[idx].Exprs - var verdict []expr.Any - var nonVerdict []expr.Any - for _, e := range existingExprs { - if _, ok := e.(*expr.Verdict); ok { + if len(match) > 0 || len(extra) > 0 || len(replaceVerdict) > 0 { + var pre, post, verdict []expr.Any + for _, e := range rules[idx].Exprs { + switch e.(type) { + case *expr.Verdict: verdict = append(verdict, e) - } else { - nonVerdict = append(nonVerdict, e) + case *expr.Log: + post = append(post, e) + default: + if len(post) > 0 { + post = append(post, e) + } else { + pre = append(pre, e) + } } } if len(replaceVerdict) > 0 { verdict = replaceVerdict } - rules[idx].Exprs = append(append(nonVerdict, extra...), verdict...) + // matches go before the log so it only fires for packets the rule matches + rules[idx].Exprs = slices.Concat(pre, match, post, extra, verdict) } } state.Rules[chain] = rules diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index e4017c3..5d70611 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -3233,3 +3233,90 @@ func TestCompile_NoLogLimitKeepsInlineLog(t *testing.T) { } } } + +func TestCompile_LogLimitSplitsAroundExtrasAndNAT(t *testing.T) { + cfg := listCfg(func(c *config.Config) { + c.Settings.LogLimit = "1/sec:10" + c.Zones["loc"] = config.Zone{Type: config.ZoneIP} + c.Interfaces = append(c.Interfaces, config.Interface{Zone: "loc", Interface: "eth1"}) + c.Rules = []config.Rule{ + {Action: config.RuleAccept, Source: "net:192.0.2.1", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, Log: "info", + Mark: "0x1", User: "0", Time: &config.TimeSpec{Start: "08:00", Stop: "17:00"}, RateLimit: "5/sec:20"}, + {Action: config.RuleLog, Source: "net:192.0.2.2", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"23"}, Log: "info", + Mark: "0x2", RateLimit: "5/sec:20"}, + {Action: config.RuleDNAT, Source: "net", Dest: "loc:198.51.100.10:80", Proto: "tcp", DPort: config.PortSpec{"8000"}, Log: "info"}, + {Action: config.RuleAccept, Source: "net:192.0.2.3", Dest: "loc", Proto: "tcp", DPort: config.PortSpec{"443"}, Log: "info"}, + } + c.SNAT = []config.SNATRule{{Action: config.SNATAddress, Source: "198.51.100.0/24", Dest: "eth0", Address: "203.0.113.7", Mark: "0x3", Log: "info"}} + }) + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + logLimit := &expr.Limit{Type: expr.LimitTypePkts, Rate: 1, Unit: expr.LimitTimeSecond, Burst: 10} + kinds := func(exprs []expr.Any) (logs, limits, marks, uids int) { + for _, e := range exprs { + switch v := e.(type) { + case *expr.Log: + logs++ + case *expr.Limit: + limits++ + case *expr.Meta: + if v.Key == expr.MetaKeyMARK && !v.SourceRegister { + marks++ + } + if v.Key == expr.MetaKeySKUID { + uids++ + } + } + } + return + } + check := func(name string, rs []ManagedRule, n, wantMarks, wantUIDs int) []expr.Any { + t.Helper() + if len(rs) != n { + t.Fatalf("%s: %d rules, want %d", name, len(rs), n) + } + l := rs[0].Exprs + if len(l) < 2 || !reflect.DeepEqual(l[len(l)-2], logLimit) { + t.Errorf("%s: log rule must end limit+log, got %#v", name, l) + } + if logs, limits, marks, uids := kinds(l); logs != 1 || limits != 1 || marks != wantMarks || uids != wantUIDs { + t.Errorf("%s: log rule logs=%d limits=%d marks=%d uids=%d, want 1/1/%d/%d", name, logs, limits, marks, uids, wantMarks, wantUIDs) + } + if n == 1 { + return nil + } + v := rs[1].Exprs + if logs, _, marks, uids := kinds(v); logs != 0 || marks != wantMarks || uids != wantUIDs { + t.Errorf("%s: action rule logs=%d marks=%d uids=%d, want 0/%d/%d", name, logs, marks, uids, wantMarks, wantUIDs) + } + return v + } + + v := check("rule:0", taggedRules(state, "input", "rule:0"), 2, 1, 1) + if _, limits, _, _ := kinds(v); limits != 1 { + t.Errorf("rule:0: action rule must keep its ratelimit, got %d limits", limits) + } + if _, ok := v[len(v)-1].(*expr.Verdict); !ok { + t.Errorf("rule:0: action rule must end in a verdict, got %T", v[len(v)-1]) + } + check("rule:1", taggedRules(state, "input", "rule:1"), 1, 1, 0) + if v := check("rule:2", taggedRules(state, "prerouting", "rule:2"), 2, 0, 0); v != nil { + if _, ok := v[len(v)-1].(*expr.NAT); !ok { + t.Errorf("rule:2: DNAT rule must end in nat, got %T", v[len(v)-1]) + } + } + check("rule:3", taggedRules(state, "forward", "rule:3"), 2, 0, 0) + var snat []ManagedRule + for _, r := range state.Rules["postrouting"] { + if logs, _, _, _ := kinds(r.Exprs); logs > 0 || strings.HasPrefix(r.Tag, "snat:0") { + snat = append(snat, r) + } + } + if v := check("snat:0", snat, 2, 1, 0); v != nil { + if _, ok := v[len(v)-1].(*expr.NAT); !ok { + t.Errorf("snat:0: rule must end in nat, got %T", v[len(v)-1]) + } + } +}