package nftables import ( "reflect" "testing" "github.com/google/nftables/expr" "golang.org/x/sys/unix" ) func TestRestoreChangeSet(t *testing.T) { accept := []expr.Any{&expr.Verdict{Kind: expr.VerdictAccept}} drop := []expr.Any{&expr.Verdict{Kind: expr.VerdictDrop}} snap := &FirewallState{Rules: map[string][]ManagedRule{ "input": { {Chain: "input", Tag: "ssh", Exprs: accept, Handle: 4}, {Chain: "input", Tag: "web", Exprs: accept, Handle: 5}, {Chain: "input", Tag: "", Exprs: drop, Handle: 6}, }, "forward": {{Chain: "forward", Tag: "fwd", Exprs: accept, Handle: 7}}, }} current := &FirewallState{Rules: map[string][]ManagedRule{ "input": { {Chain: "input", Tag: "web", Exprs: accept, Handle: 10}, {Chain: "input", Tag: "ssh", Exprs: drop, Handle: 11}, {Chain: "input", Tag: "", Exprs: drop, Handle: 12}, }, }} cs := restoreChangeSet(current, snap) var removed []uint64 for _, r := range cs.Remove { removed = append(removed, r.Handle) } if len(removed) != 2 || removed[0] != 10 || removed[1] != 11 { t.Errorf("expected managed handles [10 11] removed, untagged kept; got %v", removed) } var added []string for _, r := range cs.Add { added = append(added, r.Tag) } want := []string{"fwd", "ssh", "web"} if len(added) != len(want) { t.Fatalf("added %v, want %v", added, want) } for i := range want { if added[i] != want[i] { t.Fatalf("added %v, want %v (snapshot order per chain)", added, want) } } if !reflect.DeepEqual(cs.Add[1].Exprs, accept) { t.Error("ssh not restored to its snapshot exprs") } } func TestRestoreChangeSetEmptySnapshotRemovesAll(t *testing.T) { current := &FirewallState{Rules: map[string][]ManagedRule{ "input": {{Chain: "input", Tag: "x", Handle: 1}}, }} cs := restoreChangeSet(current, &FirewallState{Rules: map[string][]ManagedRule{}}) if len(cs.Remove) != 1 || len(cs.Add) != 0 { t.Errorf("expected 1 remove 0 add, got %d/%d", len(cs.Remove), len(cs.Add)) } } func TestDiffHelpers(t *testing.T) { h := func(name, typ string, l3 uint16, l4 uint8) Helper { return Helper{Name: name, Helper: expr.CtHelper{Name: typ, L3Proto: l3, L4Proto: l4}} } ftp := h("ftp", "ftp", unix.NFPROTO_INET, unix.IPPROTO_TCP) tftp := h("tftp", "tftp", unix.NFPROTO_INET, unix.IPPROTO_UDP) sipUDP := h("sip", "sip", unix.NFPROTO_INET, unix.IPPROTO_UDP) sipTCP := h("sip", "sip", unix.NFPROTO_INET, unix.IPPROTO_TCP) current := &FirewallState{Helpers: []Helper{tftp, ftp, sipUDP}} desired := &FirewallState{Helpers: []Helper{ftp, sipTCP}} cs := computeDiff(current, desired) if !reflect.DeepEqual(cs.RemoveHelpers, []string{"tftp", "sip"}) { t.Errorf("remove = %v", cs.RemoveHelpers) } if !reflect.DeepEqual(cs.AddHelpers, []Helper{sipTCP}) { t.Errorf("add = %v", cs.AddHelpers) } // Restore recreates every helper so kernel listing order matches the snapshot. cs = restoreChangeSet(desired, current) if !reflect.DeepEqual(cs.RemoveHelpers, []string{"ftp", "sip"}) || !reflect.DeepEqual(cs.AddHelpers, current.Helpers) { t.Errorf("restore = -%v +%v", cs.RemoveHelpers, cs.AddHelpers) } live := &FirewallState{Helpers: []Helper{h("pptp", "pptp", unix.NFPROTO_IPV4, unix.IPPROTO_TCP)}} want := &FirewallState{Helpers: []Helper{h("pptp", "pptp", unix.NFPROTO_INET, unix.IPPROTO_TCP)}} if cs := computeDiff(live, want); !cs.Empty() { t.Errorf("kernel-narrowed l3proto should not diff: %+v", cs) } if cs := computeDiff(desired, desired); !cs.Empty() { t.Errorf("identical helpers should be empty: %+v", cs) } if cs := computeDiff(&FirewallState{}, desired); cs.Empty() { t.Error("missing helpers should not be empty") } }