diff --git a/internal/nftables/cleanup.go b/internal/nftables/cleanup.go index 69c31fa..d6fe590 100644 --- a/internal/nftables/cleanup.go +++ b/internal/nftables/cleanup.go @@ -28,7 +28,7 @@ func (e *Engine) FindForeignRules() ([]ForeignRule, error) { var ourTable *nftables.Table for _, t := range tables { - if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { + if t.Name == e.cfg.Settings.TableName && t.Family == e.family() { ourTable = t break } @@ -52,7 +52,7 @@ func (e *Engine) FindForeignRules() ([]ForeignRule, error) { var foreign []ForeignRule - chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) + chains, err := e.conn.ListChainsOfTableFamily(e.family()) if err != nil { return nil, fmt.Errorf("listing chains: %w", err) } diff --git a/internal/nftables/compiler.go b/internal/nftables/compiler.go index 0eeeeef..f10db29 100644 --- a/internal/nftables/compiler.go +++ b/internal/nftables/compiler.go @@ -2049,6 +2049,13 @@ func rejectExprs(proto byte, family config.AddressFamily) []expr.Any { Code: 0, }} } + // icmpx is inet-only; ip and ip6 tables silently drop on it. + switch family { + case config.FamilyIP: + return []expr.Any{&expr.Reject{Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 3}} // port-unreachable + case config.FamilyIP6: + return []expr.Any{&expr.Reject{Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 4}} // port-unreachable + } return []expr.Any{&expr.Reject{ Type: unix.NFT_REJECT_ICMPX_UNREACH, Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH, diff --git a/internal/nftables/engine.go b/internal/nftables/engine.go index c6781c6..4373e01 100644 --- a/internal/nftables/engine.go +++ b/internal/nftables/engine.go @@ -22,9 +22,64 @@ func NewEngine(cfg *config.Config) (*Engine, error) { return &Engine{cfg: cfg, conn: conn}, nil } +var tableFamilies = map[config.AddressFamily]nftables.TableFamily{ + config.FamilyINET: nftables.TableFamilyINet, + config.FamilyIP: nftables.TableFamilyIPv4, + config.FamilyIP6: nftables.TableFamilyIPv6, +} + +func (e *Engine) family() nftables.TableFamily { + if f, ok := tableFamilies[e.cfg.Settings.AddressFamily]; ok { + return f + } + return nftables.TableFamilyINet +} + +func addressFamily(tf nftables.TableFamily) config.AddressFamily { + for f, t := range tableFamilies { + if t == tf { + return f + } + } + return config.FamilyINET +} + +// withFamily is the engine for the same table name in another address family. +func (e *Engine) withFamily(f config.AddressFamily) *Engine { + cfg := *e.cfg + cfg.Settings.AddressFamily = f + return &Engine{cfg: &cfg, conn: e.conn} +} + +// overlaps reports whether tables of families a and b filter the same traffic: +// inet covers both ip and ip6, which do not overlap each other. +func overlaps(a, b nftables.TableFamily) bool { + if a == b { + return false + } + return (a == nftables.TableFamilyINet && (b == nftables.TableFamilyIPv4 || b == nftables.TableFamilyIPv6)) || + (b == nftables.TableFamilyINet && (a == nftables.TableFamilyIPv4 || a == nftables.TableFamilyIPv6)) +} + +// staleTables are our-named tables left by a different address_family that +// would still filter the traffic this family now owns. +func (e *Engine) staleTables() ([]*nftables.Table, error) { + tables, err := e.conn.ListTables() + if err != nil { + return nil, fmt.Errorf("listing tables: %w", err) + } + var stale []*nftables.Table + for _, t := range tables { + if t.Name == e.cfg.Settings.TableName && overlaps(e.family(), t.Family) { + stale = append(stale, t) + } + } + return stale, nil +} + func (e *Engine) ensureTable() *nftables.Table { return e.conn.AddTable(&nftables.Table{ - Family: nftables.TableFamilyINet, + Family: e.family(), Name: e.cfg.Settings.TableName, }) } @@ -132,6 +187,13 @@ func (e *Engine) Apply(changes *ChangeSet) error { } func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPolicy) error { + stale, err := e.staleTables() + if err != nil { + return err + } + for _, t := range stale { + e.conn.DelTable(t) + } table := e.ensureTable() chains := e.ensureChains(table, policies) @@ -173,18 +235,24 @@ func (e *Engine) apply(changes *ChangeSet, policies map[string]nftables.ChainPol } func (e *Engine) Flush() error { - tables, err := e.conn.ListTables() + tables, err := e.staleTables() if err != nil { - return fmt.Errorf("listing tables: %w", err) + return err + } + own, err := e.findTable() + if err != nil { + return err + } + if own != nil { + tables = append(tables, own) + } + if len(tables) == 0 { + return nil } - for _, t := range tables { - if t.Name == e.cfg.Settings.TableName { - e.conn.DelTable(t) - return e.conn.Flush() - } + e.conn.DelTable(t) } - return nil + return e.conn.Flush() } func (e *Engine) findTable() (*nftables.Table, error) { @@ -193,7 +261,7 @@ func (e *Engine) findTable() (*nftables.Table, error) { return nil, fmt.Errorf("listing tables: %w", err) } for _, t := range tables { - if t.Name == e.cfg.Settings.TableName && t.Family == nftables.TableFamilyINet { + if t.Name == e.cfg.Settings.TableName && t.Family == e.family() { return t, nil } } @@ -222,7 +290,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { } } - chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) + chains, err := e.conn.ListChainsOfTableFamily(e.family()) if err != nil { return nil, fmt.Errorf("listing chains: %w", err) } @@ -252,6 +320,7 @@ func (e *Engine) readCurrentState() (*FirewallState, error) { // survives the process that took it. type Snapshot struct { Table string `json:"table"` + Family config.AddressFamily `json:"family,omitempty"` Present bool `json:"present"` Policies map[string]nftables.ChainPolicy `json:"policies,omitempty"` Rules map[string][]SnapshotRule `json:"rules,omitempty"` @@ -264,16 +333,30 @@ type SnapshotRule struct { Exprs [][]byte `json:"exprs"` } -// Snapshot captures the live tomswall table so Restore can roll back to it. +// Snapshot captures the live tomswall table so Restore can roll back to it, +// falling back to the overlapping table of another family that apply replaces. func (e *Engine) Snapshot() (*Snapshot, error) { - snap := &Snapshot{Table: e.cfg.Settings.TableName} t, err := e.findTable() - if err != nil || t == nil { - return snap, err + if err != nil { + return nil, err + } + if t == nil { + stale, err := e.staleTables() + if err != nil { + return nil, err + } + // ponytail: captures one stale table; an inet config replacing both ip and ip6 restores only the first. + if len(stale) > 0 { + return e.withFamily(addressFamily(stale[0].Family)).Snapshot() + } + } + snap := &Snapshot{Table: e.cfg.Settings.TableName, Family: addressFamily(e.family())} + if t == nil { + return snap, nil } snap.Present = true - chains, err := e.conn.ListChainsOfTableFamily(nftables.TableFamilyINet) + chains, err := e.conn.ListChainsOfTableFamily(e.family()) if err != nil { return nil, fmt.Errorf("listing chains: %w", err) } @@ -288,7 +371,7 @@ func (e *Engine) Snapshot() (*Snapshot, error) { if err != nil { return nil, err } - snap.Rules, err = encodeState(state) + snap.Rules, err = encodeState(state, byte(e.family())) if err != nil { return nil, err } @@ -302,10 +385,14 @@ func (e *Engine) Restore(s *Snapshot) error { if s.Table != e.cfg.Settings.TableName { return fmt.Errorf("snapshot is of table %q, engine manages %q", s.Table, e.cfg.Settings.TableName) } + // Snapshots predating the family field are of the inet table. + if f := addressFamily(tableFamilies[s.Family]); f != addressFamily(e.family()) { + return e.withFamily(f).Restore(s) + } if !s.Present { return e.Flush() } - want, err := decodeState(s.Rules) + want, err := decodeState(s.Rules, byte(e.family())) if err != nil { return err } diff --git a/internal/nftables/family_test.go b/internal/nftables/family_test.go new file mode 100644 index 0000000..29a6307 --- /dev/null +++ b/internal/nftables/family_test.go @@ -0,0 +1,113 @@ +package nftables + +import ( + "testing" + + "github.com/google/nftables" + "github.com/google/nftables/expr" + "github.com/mdlayher/netlink" + "golang.org/x/sys/unix" + + "git.unkin.net/unkin/tomswall/internal/config" +) + +type sentTable struct { + msg int + family nftables.TableFamily +} + +// familyEngine fakes a kernel holding tomswall tables of the given families +// and records table creations/deletions. +func familyEngine(t *testing.T, af config.AddressFamily, live ...nftables.TableFamily) (*Engine, *[]sentTable) { + var sent []sentTable + e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { + var out []netlink.Message + for _, m := range req { + switch m.Header.Type { + case nftType(unix.NFT_MSG_GETTABLE): + for _, f := range live { + attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}}) + out = append(out, netlink.Message{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append([]byte{byte(f), 0, 0, 0}, attrs...)}) + } + case nftType(unix.NFT_MSG_NEWTABLE): + sent = append(sent, sentTable{unix.NFT_MSG_NEWTABLE, nftables.TableFamily(m.Data[0])}) + case nftType(unix.NFT_MSG_DELTABLE): + sent = append(sent, sentTable{unix.NFT_MSG_DELTABLE, nftables.TableFamily(m.Data[0])}) + } + } + return out, nil + }) + e.cfg.Settings.AddressFamily = af + return e, &sent +} + +func TestApplyIPFamilyReplacesInetTable(t *testing.T) { + e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyINet, nftables.TableFamilyIPv6) + if err := e.Apply(&ChangeSet{}); err != nil { + t.Fatal(err) + } + want := []sentTable{{unix.NFT_MSG_DELTABLE, nftables.TableFamilyINet}, {unix.NFT_MSG_NEWTABLE, nftables.TableFamilyIPv4}} + if len(*sent) != len(want) || (*sent)[0] != want[0] || (*sent)[1] != want[1] { + t.Errorf("got %+v, want %+v (the ip6 table must survive)", *sent, want) + } +} + +func TestApplyInetFamilyReplacesIPTables(t *testing.T) { + e, sent := familyEngine(t, config.FamilyINET, nftables.TableFamilyIPv4, nftables.TableFamilyIPv6) + if err := e.Apply(&ChangeSet{}); err != nil { + t.Fatal(err) + } + var dels int + for _, s := range *sent { + if s.msg == unix.NFT_MSG_DELTABLE { + dels++ + } + } + if dels != 2 { + t.Errorf("want ip and ip6 tables deleted, got %+v", *sent) + } +} + +func TestFlushIPFamilyKeepsIP6Table(t *testing.T) { + e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyIPv6, nftables.TableFamilyIPv4) + if err := e.Flush(); err != nil { + t.Fatal(err) + } + if len(*sent) != 1 || (*sent)[0] != (sentTable{unix.NFT_MSG_DELTABLE, nftables.TableFamilyIPv4}) { + t.Errorf("got %+v, want only the ip table deleted", *sent) + } +} + +func TestSnapshotFallsBackToReplacedTable(t *testing.T) { + e, _ := familyEngine(t, config.FamilyIP, nftables.TableFamilyINet) + snap, err := e.Snapshot() + if err != nil { + t.Fatal(err) + } + if !snap.Present || snap.Family != config.FamilyINET { + t.Fatalf("want present inet snapshot, got %+v", snap) + } + + // Reverting to the inet snapshot drops the tried ip table. + e, sent := familyEngine(t, config.FamilyIP, nftables.TableFamilyIPv4) + if err := e.Restore(&Snapshot{Table: "tomswall", Family: config.FamilyINET, Present: true}); err != nil { + t.Fatal(err) + } + want := []sentTable{{unix.NFT_MSG_DELTABLE, nftables.TableFamilyIPv4}, {unix.NFT_MSG_NEWTABLE, nftables.TableFamilyINet}} + if len(*sent) != 2 || (*sent)[0] != want[0] || (*sent)[1] != want[1] { + t.Errorf("got %+v, want %+v", *sent, want) + } +} + +func TestRejectExprsFamily(t *testing.T) { + for af, want := range map[config.AddressFamily]expr.Reject{ + config.FamilyINET: {Type: unix.NFT_REJECT_ICMPX_UNREACH, Code: unix.NFT_REJECT_ICMPX_PORT_UNREACH}, + config.FamilyIP: {Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 3}, + config.FamilyIP6: {Type: unix.NFT_REJECT_ICMP_UNREACH, Code: 4}, + } { + got := rejectExprs(unix.IPPROTO_UDP, af)[0].(*expr.Reject) + if *got != want { + t.Errorf("%s: got %+v, want %+v", af, *got, want) + } + } +} diff --git a/internal/nftables/snapshot.go b/internal/nftables/snapshot.go index bcb9b41..54bc242 100644 --- a/internal/nftables/snapshot.go +++ b/internal/nftables/snapshot.go @@ -4,14 +4,11 @@ import ( "encoding/binary" "fmt" - "github.com/google/nftables" "github.com/google/nftables/expr" "github.com/mdlayher/netlink" "golang.org/x/sys/unix" ) -const inet = byte(nftables.TableFamilyINet) - // exprByName mirrors the expression types google/nftables can parse back from the kernel. var exprByName = map[string]func() expr.Any{ "ct": func() expr.Any { return &expr.Ct{} }, @@ -40,13 +37,13 @@ var exprByName = map[string]func() expr.Any{ "notrack": func() expr.Any { return &expr.Notrack{} }, } -func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) { +func encodeState(state *FirewallState, fam byte) (map[string][]SnapshotRule, error) { out := make(map[string][]SnapshotRule, len(state.Rules)) for chain, rules := range state.Rules { for _, r := range rules { sr := SnapshotRule{Tag: r.Tag} for _, e := range r.Exprs { - b, err := expr.Marshal(inet, e) + b, err := expr.Marshal(fam, e) if err != nil { return nil, fmt.Errorf("encoding %s rule %q: %w", chain, r.Tag, err) } @@ -58,13 +55,13 @@ func encodeState(state *FirewallState) (map[string][]SnapshotRule, error) { return out, nil } -func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) { +func decodeState(rules map[string][]SnapshotRule, fam byte) (*FirewallState, error) { state := &FirewallState{Rules: make(map[string][]ManagedRule, len(rules))} for chain, rs := range rules { for _, sr := range rs { r := ManagedRule{Chain: chain, Tag: sr.Tag} for _, b := range sr.Exprs { - e, err := decodeExpr(b) + e, err := decodeExpr(b, fam) if err != nil { return nil, fmt.Errorf("decoding %s rule %q: %w", chain, sr.Tag, err) } @@ -77,7 +74,7 @@ func decodeState(rules map[string][]SnapshotRule) (*FirewallState, error) { } // decodeExpr reverses expr.Marshal, as google/nftables does when reading rules. -func decodeExpr(b []byte) (expr.Any, error) { +func decodeExpr(b []byte, fam byte) (expr.Any, error) { ad, err := netlink.NewAttributeDecoder(b) if err != nil { return nil, err @@ -104,13 +101,13 @@ func decodeExpr(b []byte) (expr.Any, error) { if name == "notrack" { return e, nil } - if err := expr.Unmarshal(inet, data, e); err != nil { + if err := expr.Unmarshal(fam, data, e); err != nil { return nil, err } // A verdict is an immediate into the verdict register with no data. if imm, ok := e.(*expr.Immediate); ok && imm.Register == unix.NFT_REG_VERDICT && len(imm.Data) == 0 { v := &expr.Verdict{} - if err := expr.Unmarshal(inet, data, v); err != nil { + if err := expr.Unmarshal(fam, data, v); err != nil { return nil, err } return v, nil diff --git a/internal/nftables/snapshot_test.go b/internal/nftables/snapshot_test.go index babe6ba..2c45dfb 100644 --- a/internal/nftables/snapshot_test.go +++ b/internal/nftables/snapshot_test.go @@ -29,7 +29,7 @@ func TestSnapshotRulesRoundTrip(t *testing.T) { "input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}}, }} - rules, err := encodeState(state) + rules, err := encodeState(state, byte(nftables.TableFamilyINet)) if err != nil { t.Fatal(err) } @@ -45,7 +45,7 @@ func TestSnapshotRulesRoundTrip(t *testing.T) { if snap.Policies["input"] != nftables.ChainPolicyAccept { t.Errorf("policy lost: %v", snap.Policies) } - got, err := decodeState(snap.Rules) + got, err := decodeState(snap.Rules, byte(nftables.TableFamilyINet)) if err != nil { t.Fatal(err) } @@ -79,7 +79,7 @@ func TestSnapshotAndRestoreAbsentTable(t *testing.T) { for _, m := range req { sent = append(sent, m.Header.Type) if m.Header.Type == nftType(unix.NFT_MSG_GETTABLE) && tablePresent { - data := []byte{inet, 0, 0, 0} + data := []byte{byte(nftables.TableFamilyINet), 0, 0, 0} attrs, _ := netlink.MarshalAttributes([]netlink.Attribute{{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")}}) return []netlink.Message{{Header: netlink.Header{Type: nftType(unix.NFT_MSG_NEWTABLE), Sequence: m.Header.Sequence}, Data: append(data, attrs...)}}, nil } @@ -116,7 +116,7 @@ func TestRestorePresentTable(t *testing.T) { {Tag: "ssh", Exprs: []expr.Any{&expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Verdict{Kind: expr.VerdictAccept}}}, {Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}, } { - enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}) + enc, err := encodeState(&FirewallState{Rules: map[string][]ManagedRule{"input": {r}}}, byte(nftables.TableFamilyINet)) if err != nil { t.Fatal(err) } @@ -132,7 +132,7 @@ func TestRestorePresentTable(t *testing.T) { if err != nil { t.Fatal(err) } - return append([]byte{inet, 0, 0, 0}, b...) + return append([]byte{byte(nftables.TableFamilyINet), 0, 0, 0}, b...) } handle := make([]byte, 8) binary.BigEndian.PutUint64(handle, 7)