package nftables import ( "bytes" "encoding/binary" "encoding/json" "reflect" "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" ) func TestSnapshotRulesRoundTrip(t *testing.T) { exprs := []expr.Any{ &expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}}, &expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2}, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0, 22}}, &expr.Ct{Register: 1, Key: expr.CtKeySTATE}, &expr.Notrack{}, &expr.Verdict{Kind: expr.VerdictAccept}, } state := &FirewallState{Rules: map[string][]ManagedRule{ "input": {{Chain: "input", Tag: "ssh", Exprs: exprs}, {Chain: "input", Tag: "drop", Exprs: []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}}}}, }} rules, err := encodeState(state) if err != nil { t.Fatal(err) } b, err := json.Marshal(&Snapshot{Table: "tomswall", Present: true, Rules: rules, Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}}) if err != nil { t.Fatal(err) } var snap Snapshot if err := json.Unmarshal(b, &snap); err != nil { t.Fatal(err) } if snap.Policies["input"] != nftables.ChainPolicyAccept { t.Errorf("policy lost: %v", snap.Policies) } got, err := decodeState(snap.Rules) if err != nil { t.Fatal(err) } in := got.Rules["input"] if len(in) != 2 || in[0].Tag != "ssh" || in[1].Tag != "drop" { t.Fatalf("rules/order lost: %+v", in) } if !reflect.DeepEqual(in[0].Exprs, exprs) { t.Errorf("exprs changed:\n got %#v\nwant %#v", in[0].Exprs, exprs) } if _, ok := in[1].Exprs[0].(*expr.Verdict); !ok { t.Errorf("verdict decoded as %T", in[1].Exprs[0]) } } func TestEnsureChainsPolicyOverride(t *testing.T) { e := testEngine(t, nil) chains := e.ensureChains(e.ensureTable(), map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}) if *chains["input"].Policy != nftables.ChainPolicyAccept { t.Error("input policy not overridden") } if *chains["forward"].Policy != nftables.ChainPolicyDrop { t.Error("forward policy should keep its default") } } func TestSnapshotAndRestoreAbsentTable(t *testing.T) { tablePresent := false var sent []netlink.HeaderType e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { 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} 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 } } return nil, nil }) snap, err := e.Snapshot() if err != nil { t.Fatal(err) } if snap.Present || snap.Table != "tomswall" { t.Fatalf("want absent tomswall snapshot, got %+v", snap) } // The try created the table; restoring the absent snapshot deletes it. tablePresent = true sent = nil if err := e.Restore(snap); err != nil { t.Fatal(err) } deleted := false for _, ht := range sent { deleted = deleted || ht == nftType(unix.NFT_MSG_DELTABLE) } if !deleted { t.Errorf("table not deleted; sent %v", sent) } } func TestRestorePresentTable(t *testing.T) { want := []SnapshotRule{} for _, r := range []ManagedRule{ {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}}}) if err != nil { t.Fatal(err) } want = append(want, enc["input"]...) } snap := &Snapshot{Table: "tomswall", Present: true, Policies: map[string]nftables.ChainPolicy{"input": nftables.ChainPolicyAccept}, Rules: map[string][]SnapshotRule{"input": want}} // Live state: the tried ruleset left one managed rule (handle 7) in input. attrs := func(a ...netlink.Attribute) []byte { b, err := netlink.MarshalAttributes(a) if err != nil { t.Fatal(err) } return append([]byte{inet, 0, 0, 0}, b...) } handle := make([]byte, 8) binary.BigEndian.PutUint64(handle, 7) var batch []netlink.Message e := testEngine(t, func(req []netlink.Message) ([]netlink.Message, error) { if len(req) == 0 { return nil, nil } reply := func(msg int, data []byte) ([]netlink.Message, error) { return []netlink.Message{{Header: netlink.Header{Type: nftType(msg), Sequence: req[0].Header.Sequence}, Data: data}}, nil } switch req[0].Header.Type { case nftType(unix.NFT_MSG_GETTABLE): return reply(unix.NFT_MSG_NEWTABLE, attrs(netlink.Attribute{Type: unix.NFTA_TABLE_NAME, Data: []byte("tomswall\x00")})) case nftType(unix.NFT_MSG_GETCHAIN): return reply(unix.NFT_MSG_NEWCHAIN, attrs( netlink.Attribute{Type: unix.NFTA_CHAIN_TABLE, Data: []byte("tomswall\x00")}, netlink.Attribute{Type: unix.NFTA_CHAIN_NAME, Data: []byte("input\x00")})) case nftType(unix.NFT_MSG_GETRULE): return reply(unix.NFT_MSG_NEWRULE, attrs( netlink.Attribute{Type: unix.NFTA_RULE_TABLE, Data: []byte("tomswall\x00")}, netlink.Attribute{Type: unix.NFTA_RULE_CHAIN, Data: []byte("input\x00")}, netlink.Attribute{Type: unix.NFTA_RULE_HANDLE, Data: handle}, netlink.Attribute{Type: unix.NFTA_RULE_USERDATA, Data: []byte("tried")})) } batch = append(batch, req...) return nil, nil }) if err := e.Restore(snap); err != nil { t.Fatal(err) } var deleted []uint64 var added []SnapshotRule policy := map[string]uint32{} for _, m := range batch { ad, err := netlink.NewAttributeDecoder(m.Data[4:]) if err != nil { t.Fatal(err) } ad.ByteOrder = binary.BigEndian var name string var r SnapshotRule var h uint64 var pol *uint32 for ad.Next() { switch { case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_NAME: name = ad.String() case m.Header.Type == nftType(unix.NFT_MSG_NEWCHAIN) && ad.Type() == unix.NFTA_CHAIN_POLICY: v := ad.Uint32() pol = &v case m.Header.Type == nftType(unix.NFT_MSG_DELRULE) && ad.Type() == unix.NFTA_RULE_HANDLE: h = ad.Uint64() case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_USERDATA: r.Tag = string(ad.Bytes()) case m.Header.Type == nftType(unix.NFT_MSG_NEWRULE) && ad.Type() == unix.NFTA_RULE_EXPRESSIONS: ad.Nested(func(nad *netlink.AttributeDecoder) error { for nad.Next() { r.Exprs = append(r.Exprs, bytes.Clone(nad.Bytes())) } return nil }) } } switch m.Header.Type { case nftType(unix.NFT_MSG_NEWCHAIN): if pol != nil { policy[name] = *pol } case nftType(unix.NFT_MSG_DELRULE): deleted = append(deleted, h) case nftType(unix.NFT_MSG_NEWRULE): added = append(added, r) } } if !reflect.DeepEqual(deleted, []uint64{7}) { t.Errorf("deleted handles %v, want [7]", deleted) } if !reflect.DeepEqual(added, want) { t.Errorf("restored rules differ from snapshot:\n got %+v\nwant %+v", added, want) } if policy["input"] != uint32(nftables.ChainPolicyAccept) || policy["forward"] != uint32(nftables.ChainPolicyDrop) { t.Errorf("chain policies %v: input must be restored to accept, forward keep drop", policy) } } func TestRestoreRejectsOtherTable(t *testing.T) { e := testEngine(t, nil) if err := e.Restore(&Snapshot{Table: "other"}); err == nil { t.Error("expected table mismatch error") } } func nftType(msg int) netlink.HeaderType { return netlink.HeaderType(unix.NFNL_SUBSYS_NFTABLES<<8 | msg) } func testEngine(t *testing.T, dial func([]netlink.Message) ([]netlink.Message, error)) *Engine { if dial == nil { dial = func([]netlink.Message) ([]netlink.Message, error) { return nil, nil } } conn, err := nftables.New(nftables.WithTestDial(dial)) if err != nil { t.Fatal(err) } 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) } } }