114 lines
3.9 KiB
Go
114 lines
3.9 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|