diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index c557f31..0d3cadf 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -216,12 +216,12 @@ func (c *Compiler) compileConntrack(state *FirewallState) error { for i, ct := range c.cfg.Conntrack { tag := fmt.Sprintf("conntrack:%d", i) - chains := []string{"prerouting"} + chains := []string{"raw_prerouting"} switch ct.Chain { case config.ConntrackOutput: - chains = []string{"output"} + chains = []string{"raw_output"} case config.ConntrackBoth: - chains = []string{"prerouting", "output"} + chains = []string{"raw_prerouting", "raw_output"} } matches, err := l4Matches(ct.Proto, ct.DPort, nil) diff --git a/internal/nftables/compiler_test.go b/internal/nftables/compiler_test.go index 9e14e2a..d955972 100644 --- a/internal/nftables/compiler_test.go +++ b/internal/nftables/compiler_test.go @@ -593,14 +593,17 @@ func TestCompile_ConntrackNoTrack(t *testing.T) { } found := false - for _, r := range state.Rules["prerouting"] { - if r.Tag == "conntrack:0:prerouting" { + for _, r := range state.Rules["raw_prerouting"] { + if r.Tag == "conntrack:0:raw_prerouting" { found = true break } } if !found { - t.Error("no notrack rule found in prerouting chain") + t.Error("no notrack rule found in raw_prerouting chain") + } + if len(state.Rules["prerouting"]) != 0 { + t.Error("conntrack rule leaked into the nat prerouting chain") } } @@ -2177,7 +2180,7 @@ func TestCompile_ListExpansionCounts(t *testing.T) { }, "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}, + }, "raw_prerouting", "conntrack:0:raw_prerouting", 2}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -2353,7 +2356,7 @@ func TestCompile_ColonRanges(t *testing.T) { }, "postrouting", "snat:0", 2}, {"conntrack dport", func(c *config.Config) { c.Conntrack = []config.ConntrackRule{{Action: config.ConntrackNoTrack, Proto: "udp", DPort: config.PortSpec{"1024:2048"}}} - }, "prerouting", "conntrack:0:prerouting", 2}, + }, "raw_prerouting", "conntrack:0:raw_prerouting", 2}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index 77a4afe..5444970 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -69,6 +69,22 @@ func (e *Engine) ensureChains(table *nftables.Table, policies map[string]nftable Hooknum: nftables.ChainHookPrerouting, Priority: nftables.ChainPriorityNATDest, }, + "raw_prerouting": { + Name: "raw_prerouting", + Table: table, + Type: nftables.ChainTypeFilter, + Hooknum: nftables.ChainHookPrerouting, + Priority: nftables.ChainPriorityRaw, + Policy: policyPtr(nftables.ChainPolicyAccept), + }, + "raw_output": { + Name: "raw_output", + Table: table, + Type: nftables.ChainTypeFilter, + Hooknum: nftables.ChainHookOutput, + Priority: nftables.ChainPriorityRaw, + Policy: policyPtr(nftables.ChainPolicyAccept), + }, } for name, chain := range chains { diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go index 3098119..babe6ba 100644 --- a/internal/nftables/snapshot_test.go +++ b/internal/nftables/snapshot_test.go @@ -243,3 +243,17 @@ func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, e } return &Engine{cfg: &config.Config{Settings: config.Settings{TableName: "tomswall"}}, conn: conn} } + +func TestEnsureChainsRawPriority(t *testing.T) { + e := testEngine(t, nil) + chains := e.ensureChains(e.ensureTable(), nil) + for name, hook := range map[string]*nftables.ChainHook{"raw_prerouting": nftables.ChainHookPrerouting, "raw_output": nftables.ChainHookOutput} { + c, ok := chains[name] + if !ok { + t.Fatalf("%s chain not declared", name) + } + if *c.Priority != *nftables.ChainPriorityRaw || *c.Hooknum != *hook || c.Type != nftables.ChainTypeFilter || *c.Policy != nftables.ChainPolicyAccept { + t.Errorf("%s: got type %s hook %d prio %d", name, c.Type, *c.Hooknum, *c.Priority) + } + } +}