diff --git a/control.go b/control.go index 75eccef1..70176a11 100644 --- a/control.go +++ b/control.go @@ -11,6 +11,7 @@ import ( "github.com/sirupsen/logrus" "github.com/slackhq/nebula/cert" + "github.com/slackhq/nebula/firewall/events" "github.com/slackhq/nebula/header" "github.com/slackhq/nebula/overlay" ) @@ -340,6 +341,36 @@ func (c *Control) Device() overlay.Device { return c.f.inside } +// SetFirewallEventReporter installs an event reporter on the current firewall. +// Passing nil clears any installed reporter. The reporter is carried across +// firewall rule reloads. Report* methods are invoked while nebula holds +// internal locks and must be non-blocking; in particular they must not call +// back into *Control methods that touch the firewall, or deadlock will +// result. +// +// Installation is performed by shallow-copying the current *Firewall, +// setting the reporter field on the copy, and swapping the pointer under +// the conntrack lock. Every Firewall the data path sees therefore has an +// immutable reporter slot, and emit sites can read it without any +// synchronization of their own. +func (c *Control) SetFirewallEventReporter(r events.Reporter) { + old := c.f.firewall + if old == nil { + return + } + old.Conntrack.Lock() + defer old.Conntrack.Unlock() + + // Re-read under the lock in case a concurrent reload swapped in a new + // Firewall between the unlocked load above and here. Both Firewalls share + // the same Conntrack pointer in the normal (non-overflow) reload path, + // so the lock we hold is the right one for whichever we see now. + current := c.f.firewall + fw := *current + fw.reporter = r + c.f.firewall = &fw +} + func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo { chi := ControlHostInfo{ VpnAddrs: make([]netip.Addr, len(h.vpnAddrs)), diff --git a/firewall.go b/firewall.go index 93b16891..6a87e302 100644 --- a/firewall.go +++ b/firewall.go @@ -20,6 +20,7 @@ import ( "github.com/slackhq/nebula/cert" "github.com/slackhq/nebula/config" "github.com/slackhq/nebula/firewall" + "github.com/slackhq/nebula/firewall/events" ) type FirewallInterface interface { @@ -67,6 +68,14 @@ type Firewall struct { incomingMetrics firewallMetrics outgoingMetrics firewallMetrics + // reporter is the optional embedder-supplied event sink. Immutable for + // the lifetime of this Firewall; Control.SetFirewallEventReporter + // installs it by shallow-copying the Firewall under the conntrack lock + // and swapping the pointer, and reloadFirewall carries it forward. + // Read unsynchronized on the data path: the preceding Firewall-pointer + // read pins the field's value for the duration of that call. + reporter events.Reporter + l *logrus.Logger } @@ -416,23 +425,27 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table") // Drop returns an error if the packet should be dropped, explaining why. It // returns nil if the packet should not be dropped. -func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error { +func (f *Firewall) Drop(fp firewall.Packet, ctx firewall.PacketContext, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error { // Check if we spoke to this tuple, if we did then allow this packet if f.inConns(fp, h, caPool, localCache) { return nil } + peerCert := h.ConnectionState.peerCert + // Make sure remote address matches nebula certificate, and determine how to treat it if h.networks == nil { // Simple case: Certificate has one address and no unsafe networks if h.vpnAddrs[0] != fp.RemoteAddr { f.metrics(incoming).droppedRemoteAddr.Inc(1) + f.reportDrop(incoming, events.DropInvalidRemoteIP, fp, ctx, peerCert) return ErrInvalidRemoteIP } } else { nwType, ok := h.networks.Lookup(fp.RemoteAddr) if !ok { f.metrics(incoming).droppedRemoteAddr.Inc(1) + f.reportDrop(incoming, events.DropInvalidRemoteIP, fp, ctx, peerCert) return ErrInvalidRemoteIP } switch nwType { @@ -440,11 +453,13 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool * break // nothing special case NetworkTypeVPNPeer: f.metrics(incoming).droppedRemoteAddr.Inc(1) + f.reportDrop(incoming, events.DropPeerRejected, fp, ctx, peerCert) return ErrPeerRejected // reject for now, one day this may have different FW rules case NetworkTypeUnsafe: break // nothing special, one day this may have different FW rules default: f.metrics(incoming).droppedRemoteAddr.Inc(1) + f.reportDrop(incoming, events.DropUnknownNetwork, fp, ctx, peerCert) return ErrUnknownNetworkType //should never happen } } @@ -452,6 +467,7 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool * // Make sure we are supposed to be handling this local ip address if !f.routableNetworks.Contains(fp.LocalAddr) { f.metrics(incoming).droppedLocalAddr.Inc(1) + f.reportDrop(incoming, events.DropInvalidLocalIP, fp, ctx, peerCert) return ErrInvalidLocalIP } @@ -461,13 +477,14 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool * } // We now know which firewall table to check against - if !table.match(fp, incoming, h.ConnectionState.peerCert, caPool) { + if !table.match(fp, incoming, peerCert, caPool) { f.metrics(incoming).droppedNoRule.Inc(1) + f.reportDrop(incoming, events.DropNoMatchingRule, fp, ctx, peerCert) return ErrNoMatchingRule } // We always want to conntrack since it is a faster operation - f.addConn(fp, incoming) + f.addConn(fp, ctx, incoming, peerCert) return nil } @@ -486,6 +503,59 @@ func (f *Firewall) Destroy() { //TODO: clean references if/when needed } +func (f *Firewall) reportDrop(incoming bool, reason events.DropReason, fp firewall.Packet, ctx firewall.PacketContext, peerCert *cert.CachedCertificate) { + r := f.reporter + if r == nil { + return + } + r.ReportDrop(events.DropEvent{ + Incoming: incoming, + Reason: reason, + Packet: fp, + Context: ctx, + PeerCert: peerCert, + RulesVersion: f.rulesVersion, + }) +} + +func (f *Firewall) reportFlowCreate(incoming bool, fp firewall.Packet, ctx firewall.PacketContext, peerCert *cert.CachedCertificate) { + r := f.reporter + if r == nil { + return + } + r.ReportFlowCreate(events.FlowCreateEvent{ + Incoming: incoming, + Packet: fp, + Context: ctx, + PeerCert: peerCert, + RulesVersion: f.rulesVersion, + }) +} + +func (f *Firewall) reportFlowEvict(incoming bool, fp firewall.Packet, rulesVersion uint16, expired bool) { + r := f.reporter + if r == nil { + return + } + r.ReportFlowEvict(events.FlowEvictEvent{ + Incoming: incoming, + Packet: fp, + RulesVersion: rulesVersion, + Expired: expired, + }) +} + +func (f *Firewall) reportRulesReload(oldVersion, newVersion uint16) { + r := f.reporter + if r == nil { + return + } + r.ReportRulesReload(events.RulesReloadEvent{ + OldVersion: oldVersion, + NewVersion: newVersion, + }) +} + func (f *Firewall) EmitStats() { conntrack := f.Conntrack conntrack.Lock() @@ -536,7 +606,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool, WithField("oldRulesVersion", c.rulesVersion). Debugln("dropping old conntrack entry, does not match new ruleset") } + oldRulesVersion := c.rulesVersion delete(conntrack.Conns, fp) + f.reportFlowEvict(c.incoming, fp, oldRulesVersion, false) conntrack.Unlock() return false } @@ -571,7 +643,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool, return true } -func (f *Firewall) addConn(fp firewall.Packet, incoming bool) { +func (f *Firewall) addConn(fp firewall.Packet, ctx firewall.PacketContext, incoming bool, peerCert *cert.CachedCertificate) { var timeout time.Duration c := &conn{} @@ -586,7 +658,8 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) { conntrack := f.Conntrack conntrack.Lock() - if _, ok := conntrack.Conns[fp]; !ok { + _, existing := conntrack.Conns[fp] + if !existing { conntrack.TimerWheel.Advance(time.Now()) conntrack.TimerWheel.Add(fp, timeout) } @@ -597,6 +670,13 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) { c.rulesVersion = f.rulesVersion c.Expires = time.Now().Add(timeout) conntrack.Conns[fp] = c + + // Report only when this represents a genuinely new flow. Fires under the + // conntrack lock so FlowCreate/FlowEvict events stay ordered relative to + // RulesReloadEvent, which also fires under this lock. + if !existing { + f.reportFlowCreate(incoming, fp, ctx, peerCert) + } conntrack.Unlock() } @@ -620,7 +700,10 @@ func (f *Firewall) evict(p firewall.Packet) { } // This conn is done + rulesVersion := t.rulesVersion + incoming := t.incoming delete(conntrack.Conns, p) + f.reportFlowEvict(incoming, p, rulesVersion, true) } func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool { diff --git a/firewall/events/events.go b/firewall/events/events.go new file mode 100644 index 00000000..7018e977 --- /dev/null +++ b/firewall/events/events.go @@ -0,0 +1,102 @@ +// Package events defines the opt-in firewall event reporting interface. +// +// Nebula emits raw packet-level events (drops, flow creations, flow evictions, +// rule reloads) and does no aggregation, counting, batching, rule-description, +// transport, or timestamping. Embedders correlate events back to yaml rules +// out of band and capture whatever clock they need themselves. All Report* +// methods are invoked while nebula holds internal locks and must be +// non-blocking. +// +// Events are passed to Report* methods by value. Implementations must not +// take the address of a received event: doing so forces Go's escape +// analysis to move the event to the heap and costs one allocation per call. +// To forward an event, either copy its fields into the reporter's own +// pooled record or send it through a value-typed channel (chan DropEvent, +// not chan *DropEvent). +package events + +import ( + "github.com/slackhq/nebula/cert" + "github.com/slackhq/nebula/firewall" +) + +type DropReason uint8 + +const ( + DropInvalidLocalIP DropReason = iota + DropInvalidRemoteIP + DropPeerRejected + DropUnknownNetwork + DropNoMatchingRule +) + +func (r DropReason) String() string { + switch r { + case DropInvalidLocalIP: + return "invalid_local_ip" + case DropInvalidRemoteIP: + return "invalid_remote_ip" + case DropPeerRejected: + return "peer_rejected" + case DropUnknownNetwork: + return "unknown_network" + case DropNoMatchingRule: + return "no_matching_rule" + default: + return "unknown" + } +} + +// DropEvent is emitted for every packet that fails the firewall check. Drops +// are not aggregated; every drop produces one event. +type DropEvent struct { + Incoming bool + Reason DropReason + Packet firewall.Packet + Context firewall.PacketContext + PeerCert *cert.CachedCertificate + RulesVersion uint16 +} + +// FlowCreateEvent is emitted when a packet is allowed and a new conntrack +// entry is created. Subsequent packets in the same flow do not re-emit. +type FlowCreateEvent struct { + Incoming bool + Packet firewall.Packet + Context firewall.PacketContext + PeerCert *cert.CachedCertificate + RulesVersion uint16 +} + +// FlowEvictEvent is emitted when a conntrack entry is removed. Context is +// not carried: timer-wheel eviction has no packet in hand, and reload +// revalidation evicts the OLD flow rather than the triggering packet. +// RulesVersion is the version under which the flow was originally allowed, +// which may differ from the current firewall version. +type FlowEvictEvent struct { + Incoming bool + Packet firewall.Packet + RulesVersion uint16 + // Expired is true when eviction was due to conntrack timeout; false when + // the entry was removed because it failed re-validation after a reload. + Expired bool +} + +// RulesReloadEvent is emitted once after each successful firewall reload. +// Reporters that bucket state by RulesVersion should close the old bucket +// and open a new one on receipt. +type RulesReloadEvent struct { + OldVersion uint16 + NewVersion uint16 +} + +// Reporter is the embedder-supplied sink for firewall events. Implementations +// that want a timestamp should call time.Now() themselves at the top of the +// method; nebula does not provide one. See the package doc for the +// do-not-take-address rule. +type Reporter interface { + ReportDrop(DropEvent) + ReportFlowCreate(FlowCreateEvent) + ReportFlowEvict(FlowEvictEvent) + ReportRulesReload(RulesReloadEvent) +} diff --git a/firewall/packet.go b/firewall/packet.go index 2cbfb5ea..a9d4d5a3 100644 --- a/firewall/packet.go +++ b/firewall/packet.go @@ -31,6 +31,27 @@ type Packet struct { Fragment bool } +// PacketContext carries additional parsed details about a packet that are +// useful for event reporting but deliberately kept out of Packet so Packet +// can keep being used as a conntrack map key. Populated alongside Packet by +// newPacket. +// +// Fields are interpreted based on Packet.Protocol: +// - ProtoTCP: TCPFlags is meaningful; ICMPType / ICMPCode are zero +// - ProtoICMP, ProtoICMPv6: ICMPType / ICMPCode are meaningful; TCPFlags is zero +// - ProtoUDP and others: only Length is meaningful +type PacketContext struct { + // Length is the total IP packet length in bytes, including headers. + Length uint16 + // TCPFlags is the flag byte from the TCP header (bits for FIN, SYN, RST, + // PSH, ACK, URG, ECE, CWR). + TCPFlags uint8 + // ICMPType is the type field of the ICMP / ICMPv6 header. + ICMPType uint8 + // ICMPCode is the code field of the ICMP / ICMPv6 header. + ICMPCode uint8 +} + func (fp *Packet) Copy() *Packet { return &Packet{ LocalAddr: fp.LocalAddr, diff --git a/firewall_events_test.go b/firewall_events_test.go new file mode 100644 index 00000000..db6d1faa --- /dev/null +++ b/firewall_events_test.go @@ -0,0 +1,731 @@ +package nebula + +import ( + "net" + "net/netip" + "sync" + "testing" + "time" + + "github.com/gaissmai/bart" + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/slackhq/nebula/cert" + "github.com/slackhq/nebula/firewall" + "github.com/slackhq/nebula/firewall/events" + "github.com/slackhq/nebula/test" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// recordingReporter captures every event fired against it. Its methods take +// the conntrack lock implicitly (via the firewall code path that invokes +// them), so we synchronize accumulator mutations with a small mutex to keep +// the race detector happy across goroutines in case a test introduces any. +type recordingReporter struct { + mu sync.Mutex + drops []recordedDrop + creates []recordedCreate + evicts []recordedEvict + reloads []recordedReload +} + +type recordedDrop struct { + incoming bool + reason events.DropReason + remote netip.Addr + local netip.Addr + peerName string + rulesVersion uint16 + ctx firewall.PacketContext +} + +type recordedCreate struct { + incoming bool + remote netip.Addr + local netip.Addr + peerName string + rulesVersion uint16 + ctx firewall.PacketContext +} + +type recordedEvict struct { + incoming bool + remote netip.Addr + local netip.Addr + rulesVersion uint16 + expired bool +} + +type recordedReload struct { + oldVersion uint16 + newVersion uint16 +} + +func (r *recordingReporter) ReportDrop(e events.DropEvent) { + r.mu.Lock() + defer r.mu.Unlock() + name := "" + if e.PeerCert != nil && e.PeerCert.Certificate != nil { + name = e.PeerCert.Certificate.Name() + } + r.drops = append(r.drops, recordedDrop{ + incoming: e.Incoming, + reason: e.Reason, + remote: e.Packet.RemoteAddr, + local: e.Packet.LocalAddr, + peerName: name, + rulesVersion: e.RulesVersion, + ctx: e.Context, + }) +} + +func (r *recordingReporter) ReportFlowCreate(e events.FlowCreateEvent) { + r.mu.Lock() + defer r.mu.Unlock() + name := "" + if e.PeerCert != nil && e.PeerCert.Certificate != nil { + name = e.PeerCert.Certificate.Name() + } + r.creates = append(r.creates, recordedCreate{ + incoming: e.Incoming, + remote: e.Packet.RemoteAddr, + local: e.Packet.LocalAddr, + peerName: name, + rulesVersion: e.RulesVersion, + ctx: e.Context, + }) +} + +func (r *recordingReporter) ReportFlowEvict(e events.FlowEvictEvent) { + r.mu.Lock() + defer r.mu.Unlock() + r.evicts = append(r.evicts, recordedEvict{ + incoming: e.Incoming, + remote: e.Packet.RemoteAddr, + local: e.Packet.LocalAddr, + rulesVersion: e.RulesVersion, + expired: e.Expired, + }) +} + +func (r *recordingReporter) ReportRulesReload(e events.RulesReloadEvent) { + r.mu.Lock() + defer r.mu.Unlock() + r.reloads = append(r.reloads, recordedReload{ + oldVersion: e.OldVersion, + newVersion: e.NewVersion, + }) +} + +// eventFixture builds a Firewall wired to a Control plus a packet/hostinfo +// pair that a test can reuse. By default the ruleset allows the packet; +// callers mutate fw / p / h as needed before invoking Drop. +type eventFixture struct { + ctl *Control + fw *Firewall + p firewall.Packet + h *HostInfo + cp *cert.CAPool +} + +func newEventFixture(t *testing.T) *eventFixture { + t.Helper() + l := test.NewLogger() + + // myVpnNetworksTable covers our single peer address so buildNetworks takes + // the "simple case" path (h.networks stays nil); tests that want a populated + // BART table overwrite h.networks directly. + vpnNetworks := new(bart.Lite) + vpnNetworks.Insert(netip.MustParsePrefix("1.2.3.0/24")) + + // Use the same cert for "peer" and "local" endpoints, matching the + // TestFirewall_Drop fixture style: LocalAddr == RemoteAddr == peer vpn addr. + c := &dummyCert{ + name: "host1", + networks: []netip.Prefix{netip.MustParsePrefix("1.2.3.4/24")}, + groups: []string{"default-group"}, + issuer: "signer-shasum", + } + h := &HostInfo{ + ConnectionState: &ConnectionState{ + peerCert: &cert.CachedCertificate{ + Certificate: c, + InvertedGroups: map[string]struct{}{"default-group": {}}, + }, + }, + vpnAddrs: []netip.Addr{netip.MustParseAddr("1.2.3.4")}, + } + h.buildNetworks(vpnNetworks, c) + + fw := NewFirewall(l, time.Minute, time.Minute, time.Minute, c) + require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + require.NoError(t, fw.AddRule(false, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + + ctl := &Control{ + f: &Interface{firewall: fw}, + l: l, + } + + return &eventFixture{ + ctl: ctl, + fw: fw, + p: firewall.Packet{ + LocalAddr: netip.MustParseAddr("1.2.3.4"), + RemoteAddr: netip.MustParseAddr("1.2.3.4"), + LocalPort: 10, + RemotePort: 90, + Protocol: firewall.ProtoUDP, + }, + h: h, + cp: cert.NewCAPool(), + } +} + +// firewall() returns the currently-installed firewall. Needed because +// SetFirewallEventReporter replaces it via shallow-copy swap. +func (f *eventFixture) firewall() *Firewall { + return f.ctl.f.firewall +} + +func TestEvents_ReportDrop_InvalidRemoteIP(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + // Packet to an address not in the cert's networks. + f.p.RemoteAddr = netip.MustParseAddr("9.9.9.9") + assert.Equal(t, ErrInvalidRemoteIP, f.firewall().Drop(f.p, firewall.PacketContext{}, false, f.h, f.cp, nil)) + + require.Len(t, r.drops, 1) + assert.Equal(t, events.DropInvalidRemoteIP, r.drops[0].reason) + assert.False(t, r.drops[0].incoming) + assert.Equal(t, "host1", r.drops[0].peerName) + assert.Empty(t, r.creates) + assert.Empty(t, r.evicts) +} + +func TestEvents_ReportDrop_InvalidLocalIP(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + // LocalAddr outside our routable networks. + f.p.LocalAddr = netip.MustParseAddr("9.9.9.9") + assert.Equal(t, ErrInvalidLocalIP, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + + require.Len(t, r.drops, 1) + assert.Equal(t, events.DropInvalidLocalIP, r.drops[0].reason) + assert.True(t, r.drops[0].incoming) +} + +func TestEvents_ReportDrop_NoMatchingRule(t *testing.T) { + f := newEventFixture(t) + // Reset to a firewall with no matching rule. + l := test.NewLogger() + fw := NewFirewall(l, time.Minute, time.Minute, time.Minute, f.h.ConnectionState.peerCert.Certificate) + // Rule that won't match (group not in peer's groups). + require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "")) + require.NoError(t, fw.AddRule(false, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "")) + f.ctl.f.firewall = fw + + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + assert.Equal(t, ErrNoMatchingRule, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.drops, 1) + assert.Equal(t, events.DropNoMatchingRule, r.drops[0].reason) +} + +func TestEvents_ReportDrop_PeerRejected(t *testing.T) { + f := newEventFixture(t) + // Re-classify the remote as VPNPeer so it triggers DropPeerRejected. + f.h.networks = new(bart.Table[NetworkType]) + f.h.networks.Insert(netip.MustParsePrefix("1.2.3.0/24"), NetworkTypeVPNPeer) + + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + assert.Equal(t, ErrPeerRejected, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.drops, 1) + assert.Equal(t, events.DropPeerRejected, r.drops[0].reason) +} + +func TestEvents_ReportDrop_UnknownNetwork(t *testing.T) { + f := newEventFixture(t) + // Insert an unrecognized NetworkType value to hit the default branch. + f.h.networks = new(bart.Table[NetworkType]) + f.h.networks.Insert(netip.MustParsePrefix("1.2.3.0/24"), NetworkTypeUnknown) + + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + assert.Equal(t, ErrUnknownNetworkType, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.drops, 1) + assert.Equal(t, events.DropUnknownNetwork, r.drops[0].reason) +} + +func TestEvents_ReportFlowCreate_OnceOnly(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + // First allowed packet creates the conntrack entry. + require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + // Second matching packet on the same tuple is short-circuited by conntrack + // and must not fire another FlowCreate. + require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + + require.Len(t, r.creates, 1) + assert.True(t, r.creates[0].incoming) + assert.Equal(t, f.p.RemoteAddr, r.creates[0].remote) + assert.Empty(t, r.drops) +} + +func TestEvents_ReportFlowEvict_OnReloadPurge(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + // Create a flow under the current rules. + require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.creates, 1) + + // Simulate a reload that produces rules the existing flow no longer + // matches. Bump rulesVersion and replace InRules with an empty table so + // revalidation fails. + fw := f.firewall() + fw.Conntrack.Lock() + fw.rulesVersion++ + fw.InRules = newFirewallTable() + fw.Conntrack.Unlock() + + // Next packet triggers re-validation, which fails and evicts the entry. + err := fw.Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil) + assert.Equal(t, ErrNoMatchingRule, err) + + require.Len(t, r.evicts, 1) + assert.False(t, r.evicts[0].expired, "evict from reload purge is not expiration") + assert.True(t, r.evicts[0].incoming) +} + +func TestEvents_ReportFlowEvict_OnTimeout(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.creates, 1) + + // Force expiration by rewinding the entry's deadline. + fw := f.firewall() + fw.Conntrack.Lock() + c := fw.Conntrack.Conns[f.p] + require.NotNil(t, c) + c.Expires = time.Now().Add(-time.Hour) + fw.evict(f.p) + fw.Conntrack.Unlock() + + require.Len(t, r.evicts, 1) + assert.True(t, r.evicts[0].expired) +} + +func TestEvents_SetNil_Clears(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.creates, 1) + + f.ctl.SetFirewallEventReporter(nil) + resetConntrack(f.firewall()) + require.NoError(t, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + // No second create should be recorded. + assert.Len(t, r.creates, 1) +} + +func TestEvents_ReporterSurvivesSwap(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + // Simulate a reload by swapping in a fresh Firewall that carries the + // reporter forward. Mirrors what reloadFirewall does with the shared + // conntrack pointer. + l := test.NewLogger() + oldFw := f.firewall() + newFw := NewFirewall(l, time.Minute, time.Minute, time.Minute, f.h.ConnectionState.peerCert.Certificate) + require.NoError(t, newFw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + require.NoError(t, newFw.AddRule(false, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + newFw.Conntrack = oldFw.Conntrack + newFw.rulesVersion = oldFw.rulesVersion + 1 + newFw.reporter = oldFw.reporter + f.ctl.f.firewall = newFw + newFw.reportRulesReload(oldFw.rulesVersion, newFw.rulesVersion) + + require.Len(t, r.reloads, 1) + assert.Equal(t, oldFw.rulesVersion, r.reloads[0].oldVersion) + assert.Equal(t, newFw.rulesVersion, r.reloads[0].newVersion) + + // Events on the new firewall should still reach the same reporter. + require.NoError(t, newFw.Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + require.Len(t, r.creates, 1) + assert.Equal(t, newFw.rulesVersion, r.creates[0].rulesVersion) +} + +func TestEvents_InstallDoesNotMutateOldFirewall(t *testing.T) { + f := newEventFixture(t) + before := f.firewall() + + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + after := f.firewall() + assert.NotSame(t, before, after, "SetFirewallEventReporter must replace the Firewall pointer") + assert.Nil(t, before.reporter, "the pre-install Firewall must remain untouched") + assert.NotNil(t, after.reporter) +} + +// --- PacketContext parse tests -------------------------------------------- + +func mustSerialize(t *testing.T, lrs ...gopacket.SerializableLayer) []byte { + t.Helper() + buf := gopacket.NewSerializeBuffer() + opt := gopacket.SerializeOptions{ComputeChecksums: false, FixLengths: true} + require.NoError(t, gopacket.SerializeLayers(buf, opt, lrs...)) + return buf.Bytes() +} + +func TestPacketContext_IPv4_TCPFlags(t *testing.T) { + ip := &layers.IPv4{ + Version: 4, TTL: 64, Protocol: layers.IPProtocolTCP, + SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2), + } + tcp := &layers.TCP{SrcPort: 1234, DstPort: 80, SYN: true, ACK: true} + require.NoError(t, tcp.SetNetworkLayerForChecksum(ip)) + data := mustSerialize(t, ip, tcp, gopacket.Payload([]byte("hello"))) + + var fp firewall.Packet + var ctx firewall.PacketContext + require.NoError(t, newPacket(data, true, &fp, &ctx)) + + assert.Equal(t, uint8(firewall.ProtoTCP), fp.Protocol) + // SYN (0x02) + ACK (0x10) = 0x12 + assert.Equal(t, uint8(0x12), ctx.TCPFlags) + assert.Equal(t, uint16(len(data)), ctx.Length) + assert.Equal(t, uint8(0), ctx.ICMPType) + assert.Equal(t, uint8(0), ctx.ICMPCode) +} + +func TestPacketContext_IPv4_ICMPTypeCode(t *testing.T) { + ip := &layers.IPv4{ + Version: 4, TTL: 64, Protocol: layers.IPProtocolICMPv4, + SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2), + } + // Destination Unreachable, code 3 (port unreachable) + icmp := &layers.ICMPv4{ + TypeCode: layers.CreateICMPv4TypeCode(layers.ICMPv4TypeDestinationUnreachable, layers.ICMPv4CodePort), + } + data := mustSerialize(t, ip, icmp, gopacket.Payload([]byte{0, 0, 0, 0})) + + var fp firewall.Packet + var ctx firewall.PacketContext + require.NoError(t, newPacket(data, true, &fp, &ctx)) + + assert.Equal(t, uint8(firewall.ProtoICMP), fp.Protocol) + assert.Equal(t, uint8(layers.ICMPv4TypeDestinationUnreachable), ctx.ICMPType) + assert.Equal(t, uint8(layers.ICMPv4CodePort), ctx.ICMPCode) + assert.Equal(t, uint16(len(data)), ctx.Length) + assert.Equal(t, uint8(0), ctx.TCPFlags) +} + +func TestPacketContext_IPv4_UDPLengthOnly(t *testing.T) { + ip := &layers.IPv4{ + Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP, + SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2), + } + udp := &layers.UDP{SrcPort: 1234, DstPort: 53} + require.NoError(t, udp.SetNetworkLayerForChecksum(ip)) + data := mustSerialize(t, ip, udp, gopacket.Payload([]byte("query"))) + + var fp firewall.Packet + var ctx firewall.PacketContext + require.NoError(t, newPacket(data, true, &fp, &ctx)) + + assert.Equal(t, uint8(firewall.ProtoUDP), fp.Protocol) + assert.Equal(t, uint16(len(data)), ctx.Length) + assert.Zero(t, ctx.TCPFlags) + assert.Zero(t, ctx.ICMPType) + assert.Zero(t, ctx.ICMPCode) +} + +func TestPacketContext_IPv6_TCPFlags(t *testing.T) { + ip := &layers.IPv6{ + Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolTCP, + SrcIP: net.ParseIP("fd00::1"), DstIP: net.ParseIP("fd00::2"), + } + tcp := &layers.TCP{SrcPort: 1234, DstPort: 443, FIN: true, ACK: true} + require.NoError(t, tcp.SetNetworkLayerForChecksum(ip)) + data := mustSerialize(t, ip, tcp, gopacket.Payload([]byte("bye"))) + + var fp firewall.Packet + var ctx firewall.PacketContext + require.NoError(t, newPacket(data, true, &fp, &ctx)) + + assert.Equal(t, uint8(firewall.ProtoTCP), fp.Protocol) + // FIN (0x01) + ACK (0x10) = 0x11 + assert.Equal(t, uint8(0x11), ctx.TCPFlags) + assert.Equal(t, uint16(len(data)), ctx.Length) +} + +func TestPacketContext_IPv6_ICMPv6TypeCode(t *testing.T) { + ip := &layers.IPv6{ + Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolICMPv6, + SrcIP: net.ParseIP("fd00::1"), DstIP: net.ParseIP("fd00::2"), + } + icmp := &layers.ICMPv6{ + TypeCode: layers.CreateICMPv6TypeCode(layers.ICMPv6TypeDestinationUnreachable, layers.ICMPv6CodePortUnreachable), + } + require.NoError(t, icmp.SetNetworkLayerForChecksum(ip)) + data := mustSerialize(t, ip, icmp, gopacket.Payload([]byte{0, 0, 0, 0, 0, 0, 0, 0})) + + var fp firewall.Packet + var ctx firewall.PacketContext + require.NoError(t, newPacket(data, true, &fp, &ctx)) + + assert.Equal(t, uint8(firewall.ProtoICMPv6), fp.Protocol) + assert.Equal(t, uint8(layers.ICMPv6TypeDestinationUnreachable), ctx.ICMPType) + assert.Equal(t, uint8(layers.ICMPv6CodePortUnreachable), ctx.ICMPCode) + assert.Equal(t, uint16(len(data)), ctx.Length) +} + +// TestPacketContext_NilOK confirms a nil context pointer is accepted by +// newPacket (the hot path may elect not to pass one). +func TestPacketContext_NilOK(t *testing.T) { + ip := &layers.IPv4{ + Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP, + SrcIP: net.IPv4(10, 0, 0, 1), DstIP: net.IPv4(10, 0, 0, 2), + } + udp := &layers.UDP{SrcPort: 1, DstPort: 2} + require.NoError(t, udp.SetNetworkLayerForChecksum(ip)) + data := mustSerialize(t, ip, udp) + + var fp firewall.Packet + require.NoError(t, newPacket(data, true, &fp, nil)) +} + +// TestPacketContext_FlowCreateCarriesContext exercises the full Drop -> addConn +// -> ReportFlowCreate path with a realistic TCP packet and confirms the +// context makes it into the reporter. +func TestPacketContext_FlowCreateCarriesContext(t *testing.T) { + f := newEventFixture(t) + r := &recordingReporter{} + f.ctl.SetFirewallEventReporter(r) + + // Hand-construct a matching TCP packet. + ctx := firewall.PacketContext{Length: 1500, TCPFlags: 0x12} + p := f.p + p.Protocol = firewall.ProtoTCP + require.NoError(t, f.firewall().Drop(p, ctx, true, f.h, f.cp, nil)) + + require.Len(t, r.creates, 1) + assert.Equal(t, uint16(1500), r.creates[0].ctx.Length) + assert.Equal(t, uint8(0x12), r.creates[0].ctx.TCPFlags) +} + +// --- benchmarks ------------------------------------------------------------ + +// noopReporter is the cheapest possible reporter. Methods discard the event. +type noopReporter struct{} + +func (noopReporter) ReportDrop(events.DropEvent) {} +func (noopReporter) ReportFlowCreate(events.FlowCreateEvent) {} +func (noopReporter) ReportFlowEvict(events.FlowEvictEvent) {} +func (noopReporter) ReportRulesReload(events.RulesReloadEvent) { +} + +// bufferedReporter demonstrates a realistic zero-alloc reporter: each event +// is forwarded to a value-typed channel. The channel send is a memcpy into +// the channel's pre-allocated ring buffer -- no heap traffic. A background +// goroutine would drain these; the bench skips draining to keep the report +// path pure. +type bufferedReporter struct { + drops chan events.DropEvent + flows chan events.FlowCreateEvent + evicts chan events.FlowEvictEvent +} + +func newBufferedReporter(cap int) *bufferedReporter { + return &bufferedReporter{ + drops: make(chan events.DropEvent, cap), + flows: make(chan events.FlowCreateEvent, cap), + evicts: make(chan events.FlowEvictEvent, cap), + } +} + +func (r *bufferedReporter) ReportDrop(e events.DropEvent) { + select { + case r.drops <- e: + default: + } +} + +func (r *bufferedReporter) ReportFlowCreate(e events.FlowCreateEvent) { + select { + case r.flows <- e: + default: + } +} + +func (r *bufferedReporter) ReportFlowEvict(e events.FlowEvictEvent) { + select { + case r.evicts <- e: + default: + } +} + +func (r *bufferedReporter) ReportRulesReload(events.RulesReloadEvent) {} + +// pointerReporter is the anti-pattern: it takes the address of the incoming +// event struct, which forces the callee-side copy onto the heap. Kept for +// comparison so we can see the alloc cost an unwary reporter would incur. +type pointerReporter struct { + last *events.DropEvent +} + +func (r *pointerReporter) ReportDrop(e events.DropEvent) { + r.last = &e +} + +func (r *pointerReporter) ReportFlowCreate(events.FlowCreateEvent) {} +func (r *pointerReporter) ReportFlowEvict(events.FlowEvictEvent) {} +func (r *pointerReporter) ReportRulesReload(events.RulesReloadEvent) { +} + +func newBenchFixture(b *testing.B) *eventFixture { + b.Helper() + l := test.NewLogger() + + vpnNetworks := new(bart.Lite) + vpnNetworks.Insert(netip.MustParsePrefix("1.2.3.0/24")) + + c := &dummyCert{ + name: "host1", + networks: []netip.Prefix{netip.MustParsePrefix("1.2.3.4/24")}, + groups: []string{"default-group"}, + issuer: "signer-shasum", + } + h := &HostInfo{ + ConnectionState: &ConnectionState{ + peerCert: &cert.CachedCertificate{ + Certificate: c, + InvertedGroups: map[string]struct{}{"default-group": {}}, + }, + }, + vpnAddrs: []netip.Addr{netip.MustParseAddr("1.2.3.4")}, + } + h.buildNetworks(vpnNetworks, c) + + fw := NewFirewall(l, time.Minute, time.Minute, time.Minute, c) + // Inbound rule that matches our packet; outbound has no match so we can + // also benchmark the no-rule drop path. + if err := fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", ""); err != nil { + b.Fatal(err) + } + + ctl := &Control{f: &Interface{firewall: fw}, l: l} + return &eventFixture{ + ctl: ctl, + fw: fw, + p: firewall.Packet{ + LocalAddr: netip.MustParseAddr("1.2.3.4"), + RemoteAddr: netip.MustParseAddr("1.2.3.4"), + LocalPort: 10, + RemotePort: 90, + Protocol: firewall.ProtoUDP, + }, + h: h, + cp: cert.NewCAPool(), + } +} + +// BenchmarkFirewallDropPath measures the cost of Firewall.Drop on a packet +// that reaches the no-matching-rule branch (the longest drop path). Compare +// reporter shapes: +// +// nilReporter -- no reporter installed (feature cost when off) +// noopReporter -- reporter installed, methods discard args (minimum on-cost) +// bufferedReporter -- realistic zero-alloc reporter: value-typed channels +// pointerReporter -- anti-pattern that takes &composite-literal (allocates) +func BenchmarkFirewallDropPath(b *testing.B) { + run := func(b *testing.B, install func(*Control)) { + f := newBenchFixture(b) + install(f.ctl) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _ = f.firewall().Drop(f.p, firewall.PacketContext{}, false, f.h, f.cp, nil) + } + } + + b.Run("nilReporter", func(b *testing.B) { run(b, func(*Control) {}) }) + b.Run("noopReporter", func(b *testing.B) { + run(b, func(c *Control) { c.SetFirewallEventReporter(noopReporter{}) }) + }) + b.Run("bufferedReporter", func(b *testing.B) { + run(b, func(c *Control) { c.SetFirewallEventReporter(newBufferedReporter(1024)) }) + }) + b.Run("pointerReporter", func(b *testing.B) { + run(b, func(c *Control) { c.SetFirewallEventReporter(&pointerReporter{}) }) + }) +} + +// BenchmarkConntrackCreate measures Firewall.Drop for an allowed inbound +// packet on a fresh conntrack (so addConn fires each iteration). +func BenchmarkConntrackCreate(b *testing.B) { + run := func(b *testing.B, install func(*Control)) { + f := newBenchFixture(b) + install(f.ctl) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + resetConntrack(f.firewall()) + _ = f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil) + } + } + b.Run("nilReporter", func(b *testing.B) { run(b, func(*Control) {}) }) + b.Run("noopReporter", func(b *testing.B) { + run(b, func(c *Control) { c.SetFirewallEventReporter(noopReporter{}) }) + }) + b.Run("bufferedReporter", func(b *testing.B) { + run(b, func(c *Control) { c.SetFirewallEventReporter(newBufferedReporter(1024)) }) + }) +} + +// BenchmarkConntrackHit measures the hot path where a flow is already in +// conntrack and short-circuits rule evaluation. The reporter slot is checked +// only on create/evict, so this bench should show the reporter having zero +// impact regardless of install state. +func BenchmarkConntrackHit(b *testing.B) { + b.Run("nilReporter", func(b *testing.B) { + f := newBenchFixture(b) + // Prime conntrack. + require.NoError(b, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _ = f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil) + } + }) + + b.Run("noopReporter", func(b *testing.B) { + f := newBenchFixture(b) + f.ctl.SetFirewallEventReporter(noopReporter{}) + require.NoError(b, f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil)) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _ = f.firewall().Drop(f.p, firewall.PacketContext{}, true, f.h, f.cp, nil) + } + }) +} diff --git a/firewall_test.go b/firewall_test.go index a2133760..37bd8a78 100644 --- a/firewall_test.go +++ b/firewall_test.go @@ -213,44 +213,44 @@ func TestFirewall_Drop(t *testing.T) { cp := cert.NewCAPool() // Drop outbound - assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil)) + assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil)) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) // Allow outbound because conntrack - require.NoError(t, fw.Drop(p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil)) // test remote mismatch oldRemote := p.RemoteAddr p.RemoteAddr = netip.MustParseAddr("1.2.3.10") - assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrInvalidRemoteIP) p.RemoteAddr = oldRemote // ensure signer doesn't get in the way of group checks fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad")) - assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) // test caSha doesn't drop on match fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum")) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) // ensure ca name doesn't get in the way of group checks cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", "")) - assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) // test caName doesn't drop on match cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", "")) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) } func TestFirewall_DropV6(t *testing.T) { @@ -292,44 +292,44 @@ func TestFirewall_DropV6(t *testing.T) { cp := cert.NewCAPool() // Drop outbound - assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil)) + assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil)) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) // Allow outbound because conntrack - require.NoError(t, fw.Drop(p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil)) // test remote mismatch oldRemote := p.RemoteAddr p.RemoteAddr = netip.MustParseAddr("fd12::56") - assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrInvalidRemoteIP) p.RemoteAddr = oldRemote // ensure signer doesn't get in the way of group checks fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad")) - assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) // test caSha doesn't drop on match fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum")) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) // ensure ca name doesn't get in the way of group checks cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", "")) - assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) // test caName doesn't drop on match cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", "")) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) } func BenchmarkFirewallTable_match(b *testing.B) { @@ -537,10 +537,10 @@ func TestFirewall_Drop2(t *testing.T) { cp := cert.NewCAPool() // h1/c1 lacks the proper groups - require.ErrorIs(t, fw.Drop(p, true, &h1, cp, nil), ErrNoMatchingRule) + require.ErrorIs(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil), ErrNoMatchingRule) // c has the proper groups resetConntrack(fw) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) } func TestFirewall_Drop3(t *testing.T) { @@ -618,18 +618,18 @@ func TestFirewall_Drop3(t *testing.T) { cp := cert.NewCAPool() // c1 should pass because host match - require.NoError(t, fw.Drop(p, true, &h1, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil)) // c2 should pass because ca sha match resetConntrack(fw) - require.NoError(t, fw.Drop(p, true, &h2, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h2, cp, nil)) // c3 should fail because no match resetConntrack(fw) - assert.Equal(t, fw.Drop(p, true, &h3, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h3, cp, nil), ErrNoMatchingRule) // Test a remote address match fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", "")) - require.NoError(t, fw.Drop(p, true, &h1, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil)) } func TestFirewall_Drop3V6(t *testing.T) { @@ -667,7 +667,7 @@ func TestFirewall_Drop3V6(t *testing.T) { fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) cp := cert.NewCAPool() require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", "")) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) } func TestFirewall_DropConntrackReload(t *testing.T) { @@ -709,12 +709,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) { cp := cert.NewCAPool() // Drop outbound - assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, true, &h, cp, nil)) // Allow outbound because conntrack - require.NoError(t, fw.Drop(p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil)) oldFw := fw fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) @@ -723,7 +723,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) { fw.rulesVersion = oldFw.rulesVersion + 1 // Allow outbound because conntrack and new rules allow port 10 - require.NoError(t, fw.Drop(p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil)) oldFw = fw fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) @@ -732,7 +732,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) { fw.rulesVersion = oldFw.rulesVersion + 1 // Drop outbound because conntrack doesn't match new ruleset - assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) } func TestFirewall_ICMPPortBehavior(t *testing.T) { @@ -778,12 +778,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 0 p.RemotePort = 0 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil)) //now also allow outbound - require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil)) }) t.Run("nonzero ports", func(t *testing.T) { @@ -791,12 +791,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 0xabcd p.RemotePort = 0x1234 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil)) //now also allow outbound - require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil)) }) }) @@ -808,12 +808,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 0 p.RemotePort = 0 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) //now also allow outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) }) t.Run("nonzero ports, still blocked", func(t *testing.T) { @@ -821,12 +821,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 0xabcd p.RemotePort = 0x1234 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) //now also allow outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) }) t.Run("nonzero, matching ports, still blocked", func(t *testing.T) { @@ -834,12 +834,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 80 p.RemotePort = 80 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil), ErrNoMatchingRule) //now also allow outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) }) }) t.Run("Any proto, any port", func(t *testing.T) { @@ -851,12 +851,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 0 p.RemotePort = 0 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil)) //now also allow outbound - require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil)) }) t.Run("nonzero ports, allowed", func(t *testing.T) { @@ -865,15 +865,15 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) { p.LocalPort = 0xabcd p.RemotePort = 0x1234 // Drop outbound - assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + assert.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) // Allow inbound resetConntrack(fw) - require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, true, &h, cp, nil)) //now also allow outbound - require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) + require.NoError(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil)) //different ID is blocked p.RemotePort++ - require.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) + require.Equal(t, fw.Drop(*p, firewall.PacketContext{}, false, &h, cp, nil), ErrNoMatchingRule) }) }) @@ -922,7 +922,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) { Protocol: firewall.ProtoUDP, Fragment: false, } - assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP) + assert.Equal(t, fw.Drop(p, firewall.PacketContext{}, true, &h1, cp, nil), ErrInvalidRemoteIP) } func BenchmarkLookup(b *testing.B) { @@ -1336,7 +1336,7 @@ func (c *testcase) Test(t *testing.T, fw *Firewall) { t.Helper() cp := cert.NewCAPool() resetConntrack(fw) - err := fw.Drop(c.p, true, c.h, cp, nil) + err := fw.Drop(c.p, firewall.PacketContext{}, true, c.h, cp, nil) if c.err == nil { require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr) } else { diff --git a/inside.go b/inside.go index 0d53f952..a09ffade 100644 --- a/inside.go +++ b/inside.go @@ -11,8 +11,8 @@ import ( "github.com/slackhq/nebula/routing" ) -func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) { - err := newPacket(packet, false, fwPacket) +func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, fwCtx *firewall.PacketContext, nb, out []byte, q int, localCache firewall.ConntrackCache) { + err := newPacket(packet, false, fwPacket, fwCtx) if err != nil { if f.l.Level >= logrus.DebugLevel { f.l.WithField("packet", packet).Debugf("Error while validating outbound packet: %s", err) @@ -66,7 +66,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet return } - dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache) + dropReason := f.firewall.Drop(*fwPacket, *fwCtx, false, hostinfo, f.pki.GetCAPool(), localCache) if dropReason == nil { f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q) @@ -211,14 +211,15 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) { fp := &firewall.Packet{} - err := newPacket(p, false, fp) + ctx := &firewall.PacketContext{} + err := newPacket(p, false, fp, ctx) if err != nil { f.l.Warnf("error while parsing outgoing packet for firewall check; %v", err) return } // check if packet is in outbound fw rules - dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil) + dropReason := f.firewall.Drop(*fp, *ctx, false, hostinfo, f.pki.GetCAPool(), nil) if dropReason != nil { if f.l.Level >= logrus.DebugLevel { f.l.WithField("fwPacket", fp). diff --git a/interface.go b/interface.go index 6d040884..e37e4244 100644 --- a/interface.go +++ b/interface.go @@ -310,10 +310,11 @@ func (f *Interface) listenOut(i int) { plaintext := make([]byte, udp.MTU) h := &header.H{} fwPacket := &firewall.Packet{} + fwCtx := &firewall.PacketContext{} nb := make([]byte, 12, 12) err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) { - f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(f.l)) + f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, fwCtx, lhh, nb, i, ctCache.Get(f.l)) }) if err != nil && !f.closed.Load() { @@ -328,6 +329,7 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) { packet := make([]byte, mtu) out := make([]byte, mtu) fwPacket := &firewall.Packet{} + fwCtx := &firewall.PacketContext{} nb := make([]byte, 12, 12) conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.conntrackCacheTimeout) @@ -342,7 +344,7 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) { break } - f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get(f.l)) + f.consumeInsidePacket(packet[:n], fwPacket, fwCtx, nb, out, i, conntrackCache.Get(f.l)) } f.l.Debugf("overlay reader %v is done", i) @@ -400,8 +402,15 @@ func (f *Interface) reloadFirewall(c *config.C) { fw.Conntrack = conntrack } + fw.reporter = oldFw.reporter + f.firewall = fw + // Fire ReportRulesReload under the conntrack lock so the reporter cannot + // observe a FlowCreate/FlowEvict for the new rulesVersion before it + // observes the reload marker. Report* must be non-blocking. + fw.reportRulesReload(oldFw.rulesVersion, fw.rulesVersion) + oldFw.Destroy() f.l.WithField("firewallHashes", fw.GetRuleHashes()). WithField("oldFirewallHashes", oldFw.GetRuleHashes()). diff --git a/outside.go b/outside.go index eba9d887..e60e5a81 100644 --- a/outside.go +++ b/outside.go @@ -19,7 +19,7 @@ const ( minFwPacketLen = 4 ) -func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { +func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, fwCtx *firewall.PacketContext, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) { err := h.Parse(packet) if err != nil { // Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors @@ -60,7 +60,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, switch h.Subtype { case header.MessageNone: - if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) { + if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, fwCtx, nb, q, localCache) { return } case header.MessageRelay: @@ -102,7 +102,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, relay: relay, IsRelayed: true, } - f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache) + f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, fwCtx, lhf, nb, q, localCache) return case ForwardingType: // Find the target HostInfo relay object @@ -295,7 +295,10 @@ var ( ) // newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers -func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { +func newPacket(data []byte, incoming bool, fp *firewall.Packet, ctx *firewall.PacketContext) error { + if ctx != nil { + *ctx = firewall.PacketContext{} + } if len(data) < 1 { return ErrPacketTooShort } @@ -303,14 +306,14 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { version := int((data[0] >> 4) & 0x0f) switch version { case ipv4.Version: - return parseV4(data, incoming, fp) + return parseV4(data, incoming, fp, ctx) case ipv6.Version: - return parseV6(data, incoming, fp) + return parseV6(data, incoming, fp, ctx) } return ErrUnknownIPVersion } -func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { +func parseV6(data []byte, incoming bool, fp *firewall.Packet, ctx *firewall.PacketContext) error { dataLen := len(data) if dataLen < ipv6.HeaderLen { return ErrIPv6PacketTooShort @@ -355,6 +358,11 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { fp.RemotePort = 0 } fp.Fragment = false + if ctx != nil { + ctx.Length = binary.BigEndian.Uint16(data[4:6]) + uint16(ipv6.HeaderLen) + ctx.ICMPType = data[offset] + ctx.ICMPCode = data[offset+1] + } return nil case layers.IPProtocolTCP, layers.IPProtocolUDP: @@ -372,6 +380,12 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { } fp.Fragment = false + if ctx != nil { + ctx.Length = binary.BigEndian.Uint16(data[4:6]) + uint16(ipv6.HeaderLen) + if proto == layers.IPProtocolTCP && dataLen >= offset+14 { + ctx.TCPFlags = data[offset+13] + } + } return nil case layers.IPProtocolIPv6Fragment: @@ -423,7 +437,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { return ErrIPv6CouldNotFindPayload } -func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { +func parseV4(data []byte, incoming bool, fp *firewall.Packet, ctx *firewall.PacketContext) error { // Do we at least have an ipv4 header worth of data? if len(data) < ipv4.HeaderLen { return ErrIPv4PacketTooShort @@ -480,6 +494,21 @@ func parseV4(data []byte, incoming bool, fp *firewall.Packet) error { fp.RemotePort = binary.BigEndian.Uint16(data[ihl+2 : ihl+4]) //dst port } + if ctx != nil { + ctx.Length = binary.BigEndian.Uint16(data[2:4]) + if !fp.Fragment { + switch fp.Protocol { + case firewall.ProtoICMP: + ctx.ICMPType = data[ihl] + ctx.ICMPCode = data[ihl+1] + case firewall.ProtoTCP: + if len(data) >= ihl+14 { + ctx.TCPFlags = data[ihl+13] + } + } + } + } + return nil } @@ -499,7 +528,7 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet [] return out, nil } -func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) bool { +func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, fwCtx *firewall.PacketContext, nb []byte, q int, localCache firewall.ConntrackCache) bool { var err error out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb) @@ -508,7 +537,7 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out return false } - err = newPacket(out, true, fwPacket) + err = newPacket(out, true, fwPacket, fwCtx) if err != nil { hostinfo.logger(f.l).WithError(err).WithField("packet", out). Warnf("Error while validating inbound packet") @@ -521,7 +550,7 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out return false } - dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache) + dropReason := f.firewall.Drop(*fwPacket, *fwCtx, true, hostinfo, f.pki.GetCAPool(), localCache) if dropReason != nil { // NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore // This gives us a buffer to build the reject packet in diff --git a/outside_test.go b/outside_test.go index 042ccbb3..14518778 100644 --- a/outside_test.go +++ b/outside_test.go @@ -20,13 +20,13 @@ func Test_newPacket(t *testing.T) { p := &firewall.Packet{} // length fails - err := newPacket([]byte{}, true, p) + err := newPacket([]byte{}, true, p, nil) require.ErrorIs(t, err, ErrPacketTooShort) - err = newPacket([]byte{0x40}, true, p) + err = newPacket([]byte{0x40}, true, p, nil) require.ErrorIs(t, err, ErrIPv4PacketTooShort) - err = newPacket([]byte{0x60}, true, p) + err = newPacket([]byte{0x60}, true, p, nil) require.ErrorIs(t, err, ErrIPv6PacketTooShort) // length fail with ip options @@ -39,15 +39,15 @@ func Test_newPacket(t *testing.T) { } b, _ := h.Marshal() - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength) // not an ipv4 packet - err = newPacket([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p) + err = newPacket([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p, nil) require.ErrorIs(t, err, ErrUnknownIPVersion) // invalid ihl - err = newPacket([]byte{4<<4 | (8 >> 2 & 0x0f), 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p) + err = newPacket([]byte{4<<4 | (8 >> 2 & 0x0f), 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p, nil) require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength) // account for variable ip header length - incoming @@ -62,7 +62,7 @@ func Test_newPacket(t *testing.T) { b, _ = h.Marshal() b = append(b, []byte{0, 3, 0, 4}...) - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.NoError(t, err) assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) @@ -84,7 +84,7 @@ func Test_newPacket(t *testing.T) { b, _ = h.Marshal() b = append(b, []byte{0, 5, 0, 6}...) - err = newPacket(b, false, p) + err = newPacket(b, false, p, nil) require.NoError(t, err) assert.Equal(t, uint8(2), p.Protocol) @@ -114,7 +114,7 @@ func Test_newPacket_v6(t *testing.T) { err := gopacket.SerializeLayers(buffer, opt, &ip) require.NoError(t, err) - err = newPacket(buffer.Bytes(), true, p) + err = newPacket(buffer.Bytes(), true, p, nil) require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) // A v6 packet with a hop-by-hop extension @@ -148,12 +148,12 @@ func Test_newPacket_v6(t *testing.T) { // A full IPv6 header and 1 byte in the first extension, but missing // the length byte. - err = newPacket(buffer.Bytes()[:41], true, p) + err = newPacket(buffer.Bytes()[:41], true, p, nil) require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) // A full IPv6 header plus 1 full extension, but only 1 byte of the // next layer, missing length byte - err = newPacket(buffer.Bytes()[:49], true, p) + err = newPacket(buffer.Bytes()[:49], true, p, nil) require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) err = nil @@ -173,7 +173,7 @@ func Test_newPacket_v6(t *testing.T) { buffer.Clear() require.NoError(t, gopacket.SerializeLayers(buffer, opt, &ip, &icmp)) - require.Error(t, newPacket(buffer.Bytes(), true, p)) + require.Error(t, newPacket(buffer.Bytes(), true, p, nil)) buffer.Clear() echo := layers.ICMPv6Echo{ @@ -181,7 +181,7 @@ func Test_newPacket_v6(t *testing.T) { SeqNumber: 1234, } require.NoError(t, gopacket.SerializeLayers(buffer, opt, &ip, &icmp, &echo)) - require.NoError(t, newPacket(buffer.Bytes(), true, p)) + require.NoError(t, newPacket(buffer.Bytes(), true, p, nil)) assert.Equal(t, uint8(layers.IPProtocolICMPv6), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr) @@ -192,7 +192,7 @@ func Test_newPacket_v6(t *testing.T) { // A good ESP packet b := buffer.Bytes() b[6] = byte(layers.IPProtocolESP) - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.NoError(t, err) assert.Equal(t, uint8(layers.IPProtocolESP), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) @@ -204,7 +204,7 @@ func Test_newPacket_v6(t *testing.T) { // A good None packet b = buffer.Bytes() b[6] = byte(layers.IPProtocolNoNextHeader) - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.NoError(t, err) assert.Equal(t, uint8(layers.IPProtocolNoNextHeader), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) @@ -216,7 +216,7 @@ func Test_newPacket_v6(t *testing.T) { // An unknown protocol packet b = buffer.Bytes() b[6] = 255 // 255 is a reserved protocol number - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) // A good UDP packet @@ -243,7 +243,7 @@ func Test_newPacket_v6(t *testing.T) { b = buffer.Bytes() // incoming - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.NoError(t, err) assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) @@ -253,7 +253,7 @@ func Test_newPacket_v6(t *testing.T) { assert.False(t, p.Fragment) // outgoing - err = newPacket(b, false, p) + err = newPacket(b, false, p, nil) require.NoError(t, err) assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr) @@ -263,14 +263,14 @@ func Test_newPacket_v6(t *testing.T) { assert.False(t, p.Fragment) // Too short UDP packet - err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes + err = newPacket(b[:len(b)-10], false, p, nil) // pull off the last 10 bytes require.ErrorIs(t, err, ErrIPv6PacketTooShort) // A good TCP packet b[6] = byte(layers.IPProtocolTCP) // incoming - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.NoError(t, err) assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) @@ -280,7 +280,7 @@ func Test_newPacket_v6(t *testing.T) { assert.False(t, p.Fragment) // outgoing - err = newPacket(b, false, p) + err = newPacket(b, false, p, nil) require.NoError(t, err) assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr) @@ -290,7 +290,7 @@ func Test_newPacket_v6(t *testing.T) { assert.False(t, p.Fragment) // Too short TCP packet - err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes + err = newPacket(b[:len(b)-10], false, p, nil) // pull off the last 10 bytes require.ErrorIs(t, err, ErrIPv6PacketTooShort) // A good UDP packet with an AH header @@ -325,7 +325,7 @@ func Test_newPacket_v6(t *testing.T) { b = append(b, ahb...) b = append(b, udpHeader...) - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.NoError(t, err) assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) @@ -335,12 +335,12 @@ func Test_newPacket_v6(t *testing.T) { assert.False(t, p.Fragment) // Ensure buffer bounds checking during processing - err = newPacket(b[:41], true, p) + err = newPacket(b[:41], true, p, nil) require.ErrorIs(t, err, ErrIPv6PacketTooShort) // Invalid AH header b = buffer.Bytes() - err = newPacket(b, true, p) + err = newPacket(b, true, p, nil) require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) } @@ -388,7 +388,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) { firstFrag = append(firstFrag, []byte{0xde, 0xad, 0xbe, 0xef}...) // Test first fragment incoming - err = newPacket(firstFrag, true, p) + err = newPacket(firstFrag, true, p, nil) require.NoError(t, err) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr) @@ -398,7 +398,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) { assert.False(t, p.Fragment) // Test first fragment outgoing - err = newPacket(firstFrag, false, p) + err = newPacket(firstFrag, false, p, nil) require.NoError(t, err) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr) assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr) @@ -427,7 +427,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) { secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...) // Test second fragment incoming - err = newPacket(secondFrag, true, p) + err = newPacket(secondFrag, true, p, nil) require.NoError(t, err) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr) assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr) @@ -437,7 +437,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) { assert.True(t, p.Fragment) // Test second fragment outgoing - err = newPacket(secondFrag, false, p) + err = newPacket(secondFrag, false, p, nil) require.NoError(t, err) assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr) assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr) @@ -447,7 +447,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) { assert.True(t, p.Fragment) // Too short of a fragment packet - err = newPacket(secondFrag[:len(secondFrag)-10], false, p) + err = newPacket(secondFrag[:len(secondFrag)-10], false, p, nil) require.ErrorIs(t, err, ErrIPv6PacketTooShort) } @@ -529,7 +529,7 @@ func BenchmarkParseV6(b *testing.B) { b.Run("Normal", func(b *testing.B) { for i := 0; i < b.N; i++ { - if err = parseV6(normalPacket, true, fp); err != nil { + if err = parseV6(normalPacket, true, fp, nil); err != nil { b.Fatal(err) } } @@ -537,7 +537,7 @@ func BenchmarkParseV6(b *testing.B) { b.Run("FirstFragment", func(b *testing.B) { for i := 0; i < b.N; i++ { - if err = parseV6(firstFrag, true, fp); err != nil { + if err = parseV6(firstFrag, true, fp, nil); err != nil { b.Fatal(err) } } @@ -545,7 +545,7 @@ func BenchmarkParseV6(b *testing.B) { b.Run("SecondFragment", func(b *testing.B) { for i := 0; i < b.N; i++ { - if err = parseV6(secondFrag, true, fp); err != nil { + if err = parseV6(secondFrag, true, fp, nil); err != nil { b.Fatal(err) } } @@ -590,7 +590,7 @@ func BenchmarkParseV6(b *testing.B) { b.Run("200 HopByHop headers", func(b *testing.B) { for i := 0; i < b.N; i++ { - if err = parseV6(evilBytes, false, fp); err != nil { + if err = parseV6(evilBytes, false, fp, nil); err != nil { b.Fatal(err) } }