From c3049e3ed43e24f7e05396b241f45fe6db0f252f Mon Sep 17 00:00:00 2001 From: unkin-agent Date: Sun, 4 Oct 2026 15:34:11 +1100 Subject: [PATCH] honour shorewall INVALID_DISPOSITION and UNTRACKED_DISPOSITION --- internal/config/config.go | 15 +++++++++ internal/nftables/compiler.go | 39 +++++++++++++++-------- internal/nftables/compiler_test.go | 50 ++++++++++++++++++++++++++++++ internal/shorewall/convert.go | 11 +++++++ internal/shorewall/convert_test.go | 30 ++++++++++++++++++ tomswall.example.yaml | 4 +++ 6 files changed, 136 insertions(+), 13 deletions(-) diff --git a/internal/config/config.go b/internal/config/config.go index 26f1b7b..d7e50e6 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -59,6 +59,11 @@ type Settings struct { // When true, auto-generate CONTINUE policies for sub-zones to their parent zones. ImplicitContinue bool `yaml:"implicit_continue,omitempty"` + + // Verdict for ct state invalid/untracked packets; continue passes them to the rules. + // Unset: invalid drops, untracked continues. + InvalidDisposition PolicyAction `yaml:"invalid_disposition,omitempty"` + UntrackedDisposition PolicyAction `yaml:"untracked_disposition,omitempty"` } // Load reads a config file in YAML or JSON format (detected by extension). @@ -115,6 +120,16 @@ 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) } + for name, d := range map[string]PolicyAction{ + "invalid_disposition": c.Settings.InvalidDisposition, + "untracked_disposition": c.Settings.UntrackedDisposition, + } { + switch d { + case "", PolicyAccept, PolicyDrop, PolicyReject, PolicyContinue: + default: + return fmt.Errorf("unknown %s %q (use accept, drop, reject, or continue)", name, d) + } + } return nil } diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 0eeeeef..0181dba 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -70,21 +70,34 @@ func (c *Compiler) Compile() (*FirewallState, error) { } func (c *Compiler) compileConntrackFastPath(state *FirewallState) error { + invalid := c.cfg.Settings.InvalidDisposition + if invalid == "" { + invalid = config.PolicyDrop + } for _, chain := range []string{"input", "forward", "output"} { - state.Rules[chain] = append(state.Rules[chain], - ManagedRule{ + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ + Chain: chain, + Exprs: append(matchCtState(ctStateEstablished|ctStateRelated), + &expr.Verdict{Kind: expr.VerdictAccept}), + Tag: "ct:fastpath:" + chain, + }) + for _, d := range []struct { + name string + state uint32 + action config.PolicyAction + }{ + {"invalid", ctStateInvalid, invalid}, + {"untracked", ctStateUntracked, c.cfg.Settings.UntrackedDisposition}, + } { + if d.action == "" || d.action == config.PolicyContinue { + continue + } + state.Rules[chain] = append(state.Rules[chain], ManagedRule{ Chain: chain, - Exprs: append(matchCtState(ctStateEstablished|ctStateRelated), - &expr.Verdict{Kind: expr.VerdictAccept}), - Tag: "ct:fastpath:" + chain, - }, - ManagedRule{ - Chain: chain, - Exprs: append(matchCtState(ctStateInvalid), - &expr.Verdict{Kind: expr.VerdictDrop}), - Tag: "ct:invalid:" + chain, - }, - ) + Exprs: append(matchCtState(d.state), policyVerdict(d.action, c.cfg.Settings.AddressFamily)...), + Tag: "ct:" + d.name + ":" + chain, + }) + } } return nil } diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 9942833..65d6fa1 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -3105,3 +3105,53 @@ func TestSpecCount_CommaAllMatchesExpansion(t *testing.T) { } } } + +func TestCompile_Dispositions(t *testing.T) { + cases := []struct { + invalid, untracked config.PolicyAction + want map[string]expr.Any // tag prefix -> verdict expr, nil = absent + }{ + {"", "", map[string]expr.Any{"ct:invalid:": &expr.Verdict{Kind: expr.VerdictDrop}, "ct:untracked:": nil}}, + {config.PolicyContinue, config.PolicyContinue, map[string]expr.Any{"ct:invalid:": nil, "ct:untracked:": nil}}, + {config.PolicyReject, config.PolicyAccept, map[string]expr.Any{ + "ct:invalid:": rejectExprs(0, config.FamilyINET)[0], + "ct:untracked:": &expr.Verdict{Kind: expr.VerdictAccept}, + }}, + } + for _, tc := range cases { + cfg := &config.Config{ + Settings: config.Settings{ + TableName: "test", + AddressFamily: config.FamilyINET, + InvalidDisposition: tc.invalid, + UntrackedDisposition: tc.untracked, + }, + 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() + if err != nil { + t.Fatalf("Compile: %v", err) + } + for _, chain := range []string{"input", "forward", "output"} { + for prefix, want := range tc.want { + var got *ManagedRule + for i, r := range state.Rules[chain] { + if r.Tag == prefix+chain { + got = &state.Rules[chain][i] + } + } + switch { + case want == nil && got != nil: + t.Errorf("%q/%q: unexpected %s%s rule", tc.invalid, tc.untracked, prefix, chain) + case want != nil && got == nil: + t.Errorf("%q/%q: missing %s%s rule", tc.invalid, tc.untracked, prefix, chain) + case want != nil && !reflect.DeepEqual(got.Exprs[len(got.Exprs)-1], want): + t.Errorf("%q/%q: %s%s verdict = %#v, want %#v", tc.invalid, tc.untracked, prefix, chain, got.Exprs[len(got.Exprs)-1], want) + } + } + } + } +} diff --git a/internal/shorewall/convert.go b/internal/shorewall/convert.go index 7ce0f74..432009c 100644 --- a/internal/shorewall/convert.go +++ b/internal/shorewall/convert.go @@ -99,6 +99,15 @@ func convertDir(dir string, ipv6 bool) (*config.Config, error) { return cfg, nil } +// disposition maps a shorewall *_DISPOSITION value; unset means CONTINUE and A_ (audit) variants map to their base action. +func disposition(v string) config.PolicyAction { + v = strings.TrimPrefix(strings.ToLower(v), "a_") + if v == "" { + return config.PolicyContinue + } + return config.PolicyAction(v) +} + func subst(s string, params map[string]string) string { if !strings.Contains(s, "$") { return s @@ -136,6 +145,8 @@ func convertConf(dir string, cfg *config.Config, params map[string]string, ipv6 if v, ok := conf["IMPLICIT_CONTINUE"]; ok { cfg.Settings.ImplicitContinue = v == "Yes" } + cfg.Settings.InvalidDisposition = disposition(conf["INVALID_DISPOSITION"]) + cfg.Settings.UntrackedDisposition = disposition(conf["UNTRACKED_DISPOSITION"]) return nil } diff --git a/internal/shorewall/convert_test.go b/internal/shorewall/convert_test.go index 5b4b49c..dda9037 100644 --- a/internal/shorewall/convert_test.go +++ b/internal/shorewall/convert_test.go @@ -725,3 +725,33 @@ func TestIsIPv6Dir(t *testing.T) { } }) } + +func TestConvert_Dispositions(t *testing.T) { + cases := []struct { + conf string + invalid, untracked config.PolicyAction + }{ + {"IP_FORWARDING=Yes", config.PolicyContinue, config.PolicyContinue}, + {"INVALID_DISPOSITION=CONTINUE\nUNTRACKED_DISPOSITION=ACCEPT", config.PolicyContinue, config.PolicyAccept}, + {"INVALID_DISPOSITION=DROP\nUNTRACKED_DISPOSITION=A_DROP", config.PolicyDrop, config.PolicyDrop}, + {"INVALID_DISPOSITION=A_REJECT", config.PolicyReject, config.PolicyContinue}, + } + for _, tc := range cases { + dir := t.TempDir() + writeFile(t, dir, "shorewall.conf", tc.conf) + writeFile(t, dir, "zones", "fw firewall\nnet ipv4\n") + writeFile(t, dir, "interfaces", "net eth0 -\n") + writeFile(t, dir, "policy", "all all DROP\n") + cfg, err := Convert(dir) + if err != nil { + t.Fatalf("Convert(%q): %v", tc.conf, err) + } + if cfg.Settings.InvalidDisposition != tc.invalid || cfg.Settings.UntrackedDisposition != tc.untracked { + t.Errorf("%q: got invalid=%q untracked=%q, want %q/%q", tc.conf, + cfg.Settings.InvalidDisposition, cfg.Settings.UntrackedDisposition, tc.invalid, tc.untracked) + } + if err := cfg.Validate(); err != nil { + t.Errorf("%q: Validate: %v", tc.conf, err) + } + } +} diff --git a/tomswall.example.yaml b/tomswall.example.yaml index b48acf4..9792aab 100644 --- a/tomswall.example.yaml +++ b/tomswall.example.yaml @@ -8,6 +8,10 @@ settings: log_level: info table_name: tomswall implicit_continue: false + # ct state invalid/untracked verdict: accept, drop, reject, continue (pass to rules) + # defaults: invalid drop, untracked continue (migrate defaults both to continue, as shorewall) + invalid_disposition: drop + untracked_disposition: continue # Named port groups — reusable port+protocol combos referenced in rules portgroups: