diff --git a/internal/config/config.go b/internal/config/config.go index d7e50e6..6f19776 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "regexp" "strings" "gopkg.in/yaml.v3" @@ -55,7 +56,9 @@ type Settings struct { AddressFamily AddressFamily `yaml:"address_family,omitempty"` IPForwarding bool `yaml:"ip_forwarding"` LogLevel string `yaml:"log_level"` - TableName string `yaml:"table_name"` + // LogLimit rate-limits every log site, shorewall LOGLIMIT syntax rate/unit[:burst]; unset logs every hit. + LogLimit string `yaml:"log_limit,omitempty"` + TableName string `yaml:"table_name"` // When true, auto-generate CONTINUE policies for sub-zones to their parent zones. ImplicitContinue bool `yaml:"implicit_continue,omitempty"` @@ -112,6 +115,8 @@ func (c *Config) applyDefaults() { } } +var logLimitRe = regexp.MustCompile(`^[1-9][0-9]*/(sec|second|min|minute|hour|day)(:[1-9][0-9]*)?$`) + var validAddressFamilies = map[AddressFamily]bool{ FamilyINET: true, FamilyIP: true, FamilyIP6: true, } @@ -120,6 +125,9 @@ func (c *Config) validateSettings() error { if !validAddressFamilies[c.Settings.AddressFamily] { return fmt.Errorf("unknown address_family %q (use inet, ip, or ip6)", c.Settings.AddressFamily) } + if l := c.Settings.LogLimit; l != "" && !logLimitRe.MatchString(l) { + return fmt.Errorf("invalid log_limit %q (use rate/{sec|min|hour|day}[:burst]; per-source s:/d: is not supported)", l) + } for name, d := range map[string]PolicyAction{ "invalid_disposition": c.Settings.InvalidDisposition, "untracked_disposition": c.Settings.UntrackedDisposition, diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 3eb3bbd..94e7ba9 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -1069,3 +1069,15 @@ func TestValidateDispositions(t *testing.T) { checkErr(t, c.Validate(), tc.wantErr) } } + +func TestValidateLogLimit(t *testing.T) { + for v, ok := range map[string]bool{ + "": true, "1/sec": true, "1/sec:10": true, "30/minute:5": true, "2/hour": true, "1/day:1": true, + "s:1/sec:10": false, "d:1/sec": false, "1": false, "1/week": false, "0/sec": false, "1/sec:": false, + } { + c := &Config{Settings: Settings{AddressFamily: FamilyINET, LogLimit: v}} + if err := c.validateSettings(); (err == nil) != ok { + t.Errorf("log_limit %q: err = %v, want ok=%v", v, err, ok) + } + } +} diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 6bfb122..47244d4 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -65,10 +65,39 @@ func (c *Compiler) Compile() (*FirewallState, error) { return nil, fmt.Errorf("static-nat: %w", err) } c.compileMSSClamp(state) + limitLogs(state, c.cfg.Settings.LogLimit) return state, nil } +// 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. 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 + } + for chain, rules := range state.Rules { + var out []ManagedRule + for _, r := range rules { + i := slices.IndexFunc(r.Exprs, func(e expr.Any) bool { _, ok := e.(*expr.Log); return ok }) + if i < 0 { + out = append(out, r) + continue + } + logRule := r + logRule.Exprs = slices.Concat(r.Exprs[:i], parseRateLimit(spec), r.Exprs[i:i+1]) + out = append(out, logRule) + 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) + } + } + state.Rules[chain] = out + } +} + func (c *Compiler) compileConntrackFastPath(state *FirewallState) error { invalid := c.cfg.Settings.InvalidDisposition if invalid == "" { @@ -414,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)...) } @@ -439,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 65d6fa1..5d70611 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -3155,3 +3155,168 @@ func TestCompile_Dispositions(t *testing.T) { } } } + +func TestCompile_LogLimit(t *testing.T) { + cfg := &config.Config{ + Settings: config.Settings{ + TableName: "test", + AddressFamily: config.FamilyINET, + LogLimit: "1/sec:10", + }, + 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"}, + }, + Rules: []config.Rule{ + {Action: config.RuleAccept, Source: "net:192.0.2.1", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"22"}, Log: "info"}, + {Action: config.RuleLog, Source: "net:198.51.100.0/24", Dest: "fw", Proto: "tcp", DPort: config.PortSpec{"23"}, Log: "info"}, + }, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatalf("Compile() error: %v", err) + } + + want := &expr.Limit{Type: expr.LimitTypePkts, Rate: 1, Unit: expr.LimitTimeSecond, Burst: 10} + byTag := map[string][]ManagedRule{} + for _, r := range state.Rules["input"] { + byTag[r.Tag] = append(byTag[r.Tag], r) + } + for tag, n := range map[string]int{"policy:0": 2, "rule:0": 2, "rule:1": 1} { + rs := byTag[tag] + if len(rs) != n { + t.Fatalf("%s: %d rules, want %d", tag, len(rs), n) + } + logRule := rs[0].Exprs + l := len(logRule) + if l < 2 || !reflect.DeepEqual(logRule[l-2], want) { + t.Errorf("%s: want limit before log, got %#v", tag, logRule) + } + if _, ok := logRule[l-1].(*expr.Log); !ok { + t.Errorf("%s: log rule must end in log, got %T", tag, logRule[l-1]) + } + for _, e := range rs[n-1].Exprs[:len(rs[n-1].Exprs)-1] { + if _, ok := e.(*expr.Log); n == 2 && ok { + t.Errorf("%s: verdict rule still logs", tag) + } + } + } + if logs := byTag["policy:0"]; len(logs) == 2 { + if _, ok := logs[1].Exprs[len(logs[1].Exprs)-1].(*expr.Verdict); !ok { + t.Errorf("policy verdict rule must end in a verdict") + } + } +} + +func TestCompile_NoLogLimitKeepsInlineLog(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"}}, + PortGroups: make(map[string]config.PortGroup), + } + state, err := NewCompiler(cfg).Compile() + if err != nil { + t.Fatal(err) + } + for _, r := range state.Rules["input"] { + for _, e := range r.Exprs { + if _, ok := e.(*expr.Limit); ok { + t.Errorf("%s: unexpected limit without log_limit", r.Tag) + } + } + } +} + +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]) + } + } +} diff --git a/internal/shorewall/convert.go b/internal/shorewall/convert.go index 7d77edc..e9bc9ea 100644 --- a/internal/shorewall/convert.go +++ b/internal/shorewall/convert.go @@ -2,6 +2,7 @@ package shorewall import ( "fmt" + "log/slog" "strconv" "strings" @@ -141,6 +142,17 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6 } else { cfg.Settings.LogLevel = "info" } + if v := conf["LOGLIMIT"]; v != "" { + if strings.HasPrefix(v, "s:") || strings.HasPrefix(v, "d:") { + slog.Warn("shorewall: per-address LOGLIMIT is not supported, limiting each log site globally", "loglimit", v) + v = v[2:] + } + if name, rest, ok := strings.Cut(v, ":"); ok && !strings.Contains(name, "/") { + slog.Warn("shorewall: named LOGLIMIT is not supported, dropping the name", "loglimit", conf["LOGLIMIT"]) + v = rest + } + cfg.Settings.LogLimit = v + } if v, ok := conf["IP_FORWARDING"]; ok { cfg.Settings.IPForwarding = v == "Yes" || v == "On" || v == "on" || v == "Keep" } diff --git a/internal/shorewall/convert_test.go b/internal/shorewall/convert_test.go index 4e46734..5431b5a 100644 --- a/internal/shorewall/convert_test.go +++ b/internal/shorewall/convert_test.go @@ -756,3 +756,23 @@ func TestConvert_Dispositions(t *testing.T) { } } } + +func TestConvert_LogLimit(t *testing.T) { + for in, want := range map[string]string{ + `LOGLIMIT="s:1/sec:10"`: "1/sec:10", + `LOGLIMIT=2/min`: "2/min", + `LOGLIMIT=name:1/sec:5`: "1/sec:5", + `LOGLIMIT=s:name:1/sec:5`: "1/sec:5", + `LOGLIMIT=`: "", + } { + dir := minimalShorewallDir(t) + writeFile(t, dir, "shorewall.conf", "LOG_LEVEL=info\n"+in+"\n") + cfg, err := Convert(dir) + if err != nil { + t.Fatalf("%s: %v", in, err) + } + if cfg.Settings.LogLimit != want { + t.Errorf("%s: log_limit = %q, want %q", in, cfg.Settings.LogLimit, want) + } + } +} diff --git a/tomswall.example.yaml b/tomswall.example.yaml index 9792aab..9ba7eab 100644 --- a/tomswall.example.yaml +++ b/tomswall.example.yaml @@ -6,6 +6,8 @@ settings: address_family: inet ip_forwarding: true log_level: info + # rate limit for every log site (shorewall LOGLIMIT, global form): rate/{sec|min|hour|day}[:burst]; unset logs every hit + log_limit: 1/sec:10 table_name: tomswall implicit_continue: false # ct state invalid/untracked verdict: accept, drop, reject, continue (pass to rules)