mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 14:47:02 +02:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e9357ff426 |
+31
@@ -11,6 +11,7 @@ import (
|
|||||||
|
|
||||||
"github.com/sirupsen/logrus"
|
"github.com/sirupsen/logrus"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/firewall/events"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
)
|
)
|
||||||
@@ -340,6 +341,36 @@ func (c *Control) Device() overlay.Device {
|
|||||||
return c.f.inside
|
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 {
|
func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
|
||||||
chi := ControlHostInfo{
|
chi := ControlHostInfo{
|
||||||
VpnAddrs: make([]netip.Addr, len(h.vpnAddrs)),
|
VpnAddrs: make([]netip.Addr, len(h.vpnAddrs)),
|
||||||
|
|||||||
+88
-5
@@ -20,6 +20,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
|
"github.com/slackhq/nebula/firewall/events"
|
||||||
)
|
)
|
||||||
|
|
||||||
type FirewallInterface interface {
|
type FirewallInterface interface {
|
||||||
@@ -67,6 +68,14 @@ type Firewall struct {
|
|||||||
incomingMetrics firewallMetrics
|
incomingMetrics firewallMetrics
|
||||||
outgoingMetrics 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
|
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
|
// Drop returns an error if the packet should be dropped, explaining why. It
|
||||||
// returns nil if the packet should not be dropped.
|
// 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
|
// Check if we spoke to this tuple, if we did then allow this packet
|
||||||
if f.inConns(fp, h, caPool, localCache) {
|
if f.inConns(fp, h, caPool, localCache) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
peerCert := h.ConnectionState.peerCert
|
||||||
|
|
||||||
// Make sure remote address matches nebula certificate, and determine how to treat it
|
// Make sure remote address matches nebula certificate, and determine how to treat it
|
||||||
if h.networks == nil {
|
if h.networks == nil {
|
||||||
// Simple case: Certificate has one address and no unsafe networks
|
// Simple case: Certificate has one address and no unsafe networks
|
||||||
if h.vpnAddrs[0] != fp.RemoteAddr {
|
if h.vpnAddrs[0] != fp.RemoteAddr {
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropInvalidRemoteIP, fp, ctx, peerCert)
|
||||||
return ErrInvalidRemoteIP
|
return ErrInvalidRemoteIP
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
nwType, ok := h.networks.Lookup(fp.RemoteAddr)
|
nwType, ok := h.networks.Lookup(fp.RemoteAddr)
|
||||||
if !ok {
|
if !ok {
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropInvalidRemoteIP, fp, ctx, peerCert)
|
||||||
return ErrInvalidRemoteIP
|
return ErrInvalidRemoteIP
|
||||||
}
|
}
|
||||||
switch nwType {
|
switch nwType {
|
||||||
@@ -440,11 +453,13 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
|
|||||||
break // nothing special
|
break // nothing special
|
||||||
case NetworkTypeVPNPeer:
|
case NetworkTypeVPNPeer:
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
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
|
return ErrPeerRejected // reject for now, one day this may have different FW rules
|
||||||
case NetworkTypeUnsafe:
|
case NetworkTypeUnsafe:
|
||||||
break // nothing special, one day this may have different FW rules
|
break // nothing special, one day this may have different FW rules
|
||||||
default:
|
default:
|
||||||
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
f.metrics(incoming).droppedRemoteAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropUnknownNetwork, fp, ctx, peerCert)
|
||||||
return ErrUnknownNetworkType //should never happen
|
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
|
// Make sure we are supposed to be handling this local ip address
|
||||||
if !f.routableNetworks.Contains(fp.LocalAddr) {
|
if !f.routableNetworks.Contains(fp.LocalAddr) {
|
||||||
f.metrics(incoming).droppedLocalAddr.Inc(1)
|
f.metrics(incoming).droppedLocalAddr.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropInvalidLocalIP, fp, ctx, peerCert)
|
||||||
return ErrInvalidLocalIP
|
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
|
// 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.metrics(incoming).droppedNoRule.Inc(1)
|
||||||
|
f.reportDrop(incoming, events.DropNoMatchingRule, fp, ctx, peerCert)
|
||||||
return ErrNoMatchingRule
|
return ErrNoMatchingRule
|
||||||
}
|
}
|
||||||
|
|
||||||
// We always want to conntrack since it is a faster operation
|
// We always want to conntrack since it is a faster operation
|
||||||
f.addConn(fp, incoming)
|
f.addConn(fp, ctx, incoming, peerCert)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -486,6 +503,59 @@ func (f *Firewall) Destroy() {
|
|||||||
//TODO: clean references if/when needed
|
//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() {
|
func (f *Firewall) EmitStats() {
|
||||||
conntrack := f.Conntrack
|
conntrack := f.Conntrack
|
||||||
conntrack.Lock()
|
conntrack.Lock()
|
||||||
@@ -536,7 +606,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
|||||||
WithField("oldRulesVersion", c.rulesVersion).
|
WithField("oldRulesVersion", c.rulesVersion).
|
||||||
Debugln("dropping old conntrack entry, does not match new ruleset")
|
Debugln("dropping old conntrack entry, does not match new ruleset")
|
||||||
}
|
}
|
||||||
|
oldRulesVersion := c.rulesVersion
|
||||||
delete(conntrack.Conns, fp)
|
delete(conntrack.Conns, fp)
|
||||||
|
f.reportFlowEvict(c.incoming, fp, oldRulesVersion, false)
|
||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -571,7 +643,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
|||||||
return true
|
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
|
var timeout time.Duration
|
||||||
c := &conn{}
|
c := &conn{}
|
||||||
|
|
||||||
@@ -586,7 +658,8 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
|||||||
|
|
||||||
conntrack := f.Conntrack
|
conntrack := f.Conntrack
|
||||||
conntrack.Lock()
|
conntrack.Lock()
|
||||||
if _, ok := conntrack.Conns[fp]; !ok {
|
_, existing := conntrack.Conns[fp]
|
||||||
|
if !existing {
|
||||||
conntrack.TimerWheel.Advance(time.Now())
|
conntrack.TimerWheel.Advance(time.Now())
|
||||||
conntrack.TimerWheel.Add(fp, timeout)
|
conntrack.TimerWheel.Add(fp, timeout)
|
||||||
}
|
}
|
||||||
@@ -597,6 +670,13 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
|||||||
c.rulesVersion = f.rulesVersion
|
c.rulesVersion = f.rulesVersion
|
||||||
c.Expires = time.Now().Add(timeout)
|
c.Expires = time.Now().Add(timeout)
|
||||||
conntrack.Conns[fp] = c
|
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()
|
conntrack.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -620,7 +700,10 @@ func (f *Firewall) evict(p firewall.Packet) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// This conn is done
|
// This conn is done
|
||||||
|
rulesVersion := t.rulesVersion
|
||||||
|
incoming := t.incoming
|
||||||
delete(conntrack.Conns, p)
|
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 {
|
func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -31,6 +31,27 @@ type Packet struct {
|
|||||||
Fragment bool
|
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 {
|
func (fp *Packet) Copy() *Packet {
|
||||||
return &Packet{
|
return &Packet{
|
||||||
LocalAddr: fp.LocalAddr,
|
LocalAddr: fp.LocalAddr,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+52
-52
@@ -213,44 +213,44 @@ func TestFirewall_Drop(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
// 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
|
// test remote mismatch
|
||||||
oldRemote := p.RemoteAddr
|
oldRemote := p.RemoteAddr
|
||||||
p.RemoteAddr = netip.MustParseAddr("1.2.3.10")
|
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
|
p.RemoteAddr = oldRemote
|
||||||
|
|
||||||
// ensure signer doesn't get in the way of group checks
|
// ensure signer doesn't get in the way of group checks
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
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
|
// test caSha doesn't drop on match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
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
|
// ensure ca name doesn't get in the way of group checks
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
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
|
// test caName doesn't drop on match
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "ca-good-bad", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
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) {
|
func TestFirewall_DropV6(t *testing.T) {
|
||||||
@@ -292,44 +292,44 @@ func TestFirewall_DropV6(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
// 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
|
// test remote mismatch
|
||||||
oldRemote := p.RemoteAddr
|
oldRemote := p.RemoteAddr
|
||||||
p.RemoteAddr = netip.MustParseAddr("fd12::56")
|
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
|
p.RemoteAddr = oldRemote
|
||||||
|
|
||||||
// ensure signer doesn't get in the way of group checks
|
// ensure signer doesn't get in the way of group checks
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
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
|
// test caSha doesn't drop on match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
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
|
// ensure ca name doesn't get in the way of group checks
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
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
|
// test caName doesn't drop on match
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
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{"nope"}, "", "", "", "ca-good-bad", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
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) {
|
func BenchmarkFirewallTable_match(b *testing.B) {
|
||||||
@@ -537,10 +537,10 @@ func TestFirewall_Drop2(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// h1/c1 lacks the proper groups
|
// 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
|
// c has the proper groups
|
||||||
resetConntrack(fw)
|
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) {
|
func TestFirewall_Drop3(t *testing.T) {
|
||||||
@@ -618,18 +618,18 @@ func TestFirewall_Drop3(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// c1 should pass because host match
|
// 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
|
// c2 should pass because ca sha match
|
||||||
resetConntrack(fw)
|
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
|
// c3 should fail because no match
|
||||||
resetConntrack(fw)
|
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
|
// Test a remote address match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
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.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) {
|
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)
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", ""))
|
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) {
|
func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||||
@@ -709,12 +709,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
// 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
|
oldFw := fw
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
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
|
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
|
||||||
// Allow outbound because conntrack and new rules allow port 10
|
// 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
|
oldFw = fw
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
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
|
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
|
||||||
// Drop outbound because conntrack doesn't match new ruleset
|
// 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) {
|
func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||||
@@ -778,12 +778,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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) {
|
t.Run("nonzero ports", func(t *testing.T) {
|
||||||
@@ -791,12 +791,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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) {
|
t.Run("nonzero ports, still blocked", func(t *testing.T) {
|
||||||
@@ -821,12 +821,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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) {
|
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.LocalPort = 80
|
||||||
p.RemotePort = 80
|
p.RemotePort = 80
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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) {
|
t.Run("Any proto, any port", func(t *testing.T) {
|
||||||
@@ -851,12 +851,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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) {
|
t.Run("nonzero ports, allowed", func(t *testing.T) {
|
||||||
@@ -865,15 +865,15 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// 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
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
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
|
//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
|
//different ID is blocked
|
||||||
p.RemotePort++
|
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,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
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) {
|
func BenchmarkLookup(b *testing.B) {
|
||||||
@@ -1336,7 +1336,7 @@ func (c *testcase) Test(t *testing.T, fw *Firewall) {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
resetConntrack(fw)
|
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 {
|
if c.err == nil {
|
||||||
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -11,8 +11,8 @@ import (
|
|||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
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)
|
err := newPacket(packet, false, fwPacket, fwCtx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Level >= logrus.DebugLevel {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
f.l.WithField("packet", packet).Debugf("Error while validating outbound packet: %s", err)
|
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
|
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 {
|
if dropReason == nil {
|
||||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
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) {
|
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
||||||
fp := &firewall.Packet{}
|
fp := &firewall.Packet{}
|
||||||
err := newPacket(p, false, fp)
|
ctx := &firewall.PacketContext{}
|
||||||
|
err := newPacket(p, false, fp, ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Warnf("error while parsing outgoing packet for firewall check; %v", err)
|
f.l.Warnf("error while parsing outgoing packet for firewall check; %v", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// check if packet is in outbound fw rules
|
// 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 dropReason != nil {
|
||||||
if f.l.Level >= logrus.DebugLevel {
|
if f.l.Level >= logrus.DebugLevel {
|
||||||
f.l.WithField("fwPacket", fp).
|
f.l.WithField("fwPacket", fp).
|
||||||
|
|||||||
+11
-2
@@ -310,10 +310,11 @@ func (f *Interface) listenOut(i int) {
|
|||||||
plaintext := make([]byte, udp.MTU)
|
plaintext := make([]byte, udp.MTU)
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
|
fwCtx := &firewall.PacketContext{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
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() {
|
if err != nil && !f.closed.Load() {
|
||||||
@@ -328,6 +329,7 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
|||||||
packet := make([]byte, mtu)
|
packet := make([]byte, mtu)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
|
fwCtx := &firewall.PacketContext{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.conntrackCacheTimeout)
|
||||||
@@ -342,7 +344,7 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
|||||||
break
|
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)
|
f.l.Debugf("overlay reader %v is done", i)
|
||||||
@@ -400,8 +402,15 @@ func (f *Interface) reloadFirewall(c *config.C) {
|
|||||||
fw.Conntrack = conntrack
|
fw.Conntrack = conntrack
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fw.reporter = oldFw.reporter
|
||||||
|
|
||||||
f.firewall = fw
|
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()
|
oldFw.Destroy()
|
||||||
f.l.WithField("firewallHashes", fw.GetRuleHashes()).
|
f.l.WithField("firewallHashes", fw.GetRuleHashes()).
|
||||||
WithField("oldFirewallHashes", oldFw.GetRuleHashes()).
|
WithField("oldFirewallHashes", oldFw.GetRuleHashes()).
|
||||||
|
|||||||
+40
-11
@@ -19,7 +19,7 @@ const (
|
|||||||
minFwPacketLen = 4
|
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)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// 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 {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
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
|
return
|
||||||
}
|
}
|
||||||
case header.MessageRelay:
|
case header.MessageRelay:
|
||||||
@@ -102,7 +102,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
relay: relay,
|
relay: relay,
|
||||||
IsRelayed: true,
|
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
|
return
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Find the target HostInfo relay object
|
// 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
|
// 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 {
|
if len(data) < 1 {
|
||||||
return ErrPacketTooShort
|
return ErrPacketTooShort
|
||||||
}
|
}
|
||||||
@@ -303,14 +306,14 @@ func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
version := int((data[0] >> 4) & 0x0f)
|
version := int((data[0] >> 4) & 0x0f)
|
||||||
switch version {
|
switch version {
|
||||||
case ipv4.Version:
|
case ipv4.Version:
|
||||||
return parseV4(data, incoming, fp)
|
return parseV4(data, incoming, fp, ctx)
|
||||||
case ipv6.Version:
|
case ipv6.Version:
|
||||||
return parseV6(data, incoming, fp)
|
return parseV6(data, incoming, fp, ctx)
|
||||||
}
|
}
|
||||||
return ErrUnknownIPVersion
|
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)
|
dataLen := len(data)
|
||||||
if dataLen < ipv6.HeaderLen {
|
if dataLen < ipv6.HeaderLen {
|
||||||
return ErrIPv6PacketTooShort
|
return ErrIPv6PacketTooShort
|
||||||
@@ -355,6 +358,11 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
fp.RemotePort = 0
|
fp.RemotePort = 0
|
||||||
}
|
}
|
||||||
fp.Fragment = false
|
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
|
return nil
|
||||||
|
|
||||||
case layers.IPProtocolTCP, layers.IPProtocolUDP:
|
case layers.IPProtocolTCP, layers.IPProtocolUDP:
|
||||||
@@ -372,6 +380,12 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fp.Fragment = false
|
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
|
return nil
|
||||||
|
|
||||||
case layers.IPProtocolIPv6Fragment:
|
case layers.IPProtocolIPv6Fragment:
|
||||||
@@ -423,7 +437,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
|||||||
return ErrIPv6CouldNotFindPayload
|
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?
|
// Do we at least have an ipv4 header worth of data?
|
||||||
if len(data) < ipv4.HeaderLen {
|
if len(data) < ipv4.HeaderLen {
|
||||||
return ErrIPv4PacketTooShort
|
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
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -499,7 +528,7 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []
|
|||||||
return out, nil
|
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
|
var err error
|
||||||
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
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
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
err = newPacket(out, true, fwPacket)
|
err = newPacket(out, true, fwPacket, fwCtx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).WithError(err).WithField("packet", out).
|
hostinfo.logger(f.l).WithError(err).WithField("packet", out).
|
||||||
Warnf("Error while validating inbound packet")
|
Warnf("Error while validating inbound packet")
|
||||||
@@ -521,7 +550,7 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out
|
|||||||
return false
|
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 {
|
if dropReason != nil {
|
||||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
// 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
|
// This gives us a buffer to build the reject packet in
|
||||||
|
|||||||
+34
-34
@@ -20,13 +20,13 @@ func Test_newPacket(t *testing.T) {
|
|||||||
p := &firewall.Packet{}
|
p := &firewall.Packet{}
|
||||||
|
|
||||||
// length fails
|
// length fails
|
||||||
err := newPacket([]byte{}, true, p)
|
err := newPacket([]byte{}, true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrPacketTooShort)
|
require.ErrorIs(t, err, ErrPacketTooShort)
|
||||||
|
|
||||||
err = newPacket([]byte{0x40}, true, p)
|
err = newPacket([]byte{0x40}, true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv4PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv4PacketTooShort)
|
||||||
|
|
||||||
err = newPacket([]byte{0x60}, true, p)
|
err = newPacket([]byte{0x60}, true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// length fail with ip options
|
// length fail with ip options
|
||||||
@@ -39,15 +39,15 @@ func Test_newPacket(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
b, _ := h.Marshal()
|
b, _ := h.Marshal()
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
||||||
|
|
||||||
// not an ipv4 packet
|
// 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)
|
require.ErrorIs(t, err, ErrUnknownIPVersion)
|
||||||
|
|
||||||
// invalid ihl
|
// 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)
|
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
||||||
|
|
||||||
// account for variable ip header length - incoming
|
// account for variable ip header length - incoming
|
||||||
@@ -62,7 +62,7 @@ func Test_newPacket(t *testing.T) {
|
|||||||
|
|
||||||
b, _ = h.Marshal()
|
b, _ = h.Marshal()
|
||||||
b = append(b, []byte{0, 3, 0, 4}...)
|
b = append(b, []byte{0, 3, 0, 4}...)
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
@@ -84,7 +84,7 @@ func Test_newPacket(t *testing.T) {
|
|||||||
|
|
||||||
b, _ = h.Marshal()
|
b, _ = h.Marshal()
|
||||||
b = append(b, []byte{0, 5, 0, 6}...)
|
b = append(b, []byte{0, 5, 0, 6}...)
|
||||||
err = newPacket(b, false, p)
|
err = newPacket(b, false, p, nil)
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(2), p.Protocol)
|
assert.Equal(t, uint8(2), p.Protocol)
|
||||||
@@ -114,7 +114,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
err := gopacket.SerializeLayers(buffer, opt, &ip)
|
err := gopacket.SerializeLayers(buffer, opt, &ip)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
err = newPacket(buffer.Bytes(), true, p)
|
err = newPacket(buffer.Bytes(), true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A v6 packet with a hop-by-hop extension
|
// 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
|
// A full IPv6 header and 1 byte in the first extension, but missing
|
||||||
// the length byte.
|
// the length byte.
|
||||||
err = newPacket(buffer.Bytes()[:41], true, p)
|
err = newPacket(buffer.Bytes()[:41], true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
||||||
// next layer, missing length byte
|
// 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)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
err = nil
|
err = nil
|
||||||
|
|
||||||
@@ -173,7 +173,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
|
|
||||||
buffer.Clear()
|
buffer.Clear()
|
||||||
require.NoError(t, gopacket.SerializeLayers(buffer, opt, &ip, &icmp))
|
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()
|
buffer.Clear()
|
||||||
echo := layers.ICMPv6Echo{
|
echo := layers.ICMPv6Echo{
|
||||||
@@ -181,7 +181,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
SeqNumber: 1234,
|
SeqNumber: 1234,
|
||||||
}
|
}
|
||||||
require.NoError(t, gopacket.SerializeLayers(buffer, opt, &ip, &icmp, &echo))
|
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, uint8(layers.IPProtocolICMPv6), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
||||||
@@ -192,7 +192,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// A good ESP packet
|
// A good ESP packet
|
||||||
b := buffer.Bytes()
|
b := buffer.Bytes()
|
||||||
b[6] = byte(layers.IPProtocolESP)
|
b[6] = byte(layers.IPProtocolESP)
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(layers.IPProtocolESP), p.Protocol)
|
assert.Equal(t, uint8(layers.IPProtocolESP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
@@ -204,7 +204,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// A good None packet
|
// A good None packet
|
||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
b[6] = byte(layers.IPProtocolNoNextHeader)
|
b[6] = byte(layers.IPProtocolNoNextHeader)
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(layers.IPProtocolNoNextHeader), p.Protocol)
|
assert.Equal(t, uint8(layers.IPProtocolNoNextHeader), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
@@ -216,7 +216,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// An unknown protocol packet
|
// An unknown protocol packet
|
||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
b[6] = 255 // 255 is a reserved protocol number
|
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)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A good UDP packet
|
// A good UDP packet
|
||||||
@@ -243,7 +243,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
|
|
||||||
// incoming
|
// incoming
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
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)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// outgoing
|
// outgoing
|
||||||
err = newPacket(b, false, p)
|
err = newPacket(b, false, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
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)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// Too short UDP packet
|
// 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)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// A good TCP packet
|
// A good TCP packet
|
||||||
b[6] = byte(layers.IPProtocolTCP)
|
b[6] = byte(layers.IPProtocolTCP)
|
||||||
|
|
||||||
// incoming
|
// incoming
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
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)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// outgoing
|
// outgoing
|
||||||
err = newPacket(b, false, p)
|
err = newPacket(b, false, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
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)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// Too short TCP packet
|
// 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)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// A good UDP packet with an AH header
|
// 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, ahb...)
|
||||||
b = append(b, udpHeader...)
|
b = append(b, udpHeader...)
|
||||||
|
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
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)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// Ensure buffer bounds checking during processing
|
// Ensure buffer bounds checking during processing
|
||||||
err = newPacket(b[:41], true, p)
|
err = newPacket(b[:41], true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// Invalid AH header
|
// Invalid AH header
|
||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p, nil)
|
||||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -388,7 +388,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) {
|
|||||||
firstFrag = append(firstFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
firstFrag = append(firstFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
||||||
|
|
||||||
// Test first fragment incoming
|
// Test first fragment incoming
|
||||||
err = newPacket(firstFrag, true, p)
|
err = newPacket(firstFrag, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
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)
|
assert.False(t, p.Fragment)
|
||||||
|
|
||||||
// Test first fragment outgoing
|
// Test first fragment outgoing
|
||||||
err = newPacket(firstFrag, false, p)
|
err = newPacket(firstFrag, false, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
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}...)
|
secondFrag = append(secondFrag, []byte{0xde, 0xad, 0xbe, 0xef}...)
|
||||||
|
|
||||||
// Test second fragment incoming
|
// Test second fragment incoming
|
||||||
err = newPacket(secondFrag, true, p)
|
err = newPacket(secondFrag, true, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.RemoteAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.LocalAddr)
|
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)
|
assert.True(t, p.Fragment)
|
||||||
|
|
||||||
// Test second fragment outgoing
|
// Test second fragment outgoing
|
||||||
err = newPacket(secondFrag, false, p)
|
err = newPacket(secondFrag, false, p, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
assert.Equal(t, netip.MustParseAddr("ff02::2"), p.LocalAddr)
|
||||||
assert.Equal(t, netip.MustParseAddr("ff02::1"), p.RemoteAddr)
|
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)
|
assert.True(t, p.Fragment)
|
||||||
|
|
||||||
// Too short of a fragment packet
|
// 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)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -529,7 +529,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("Normal", func(b *testing.B) {
|
b.Run("Normal", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
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)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -537,7 +537,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("FirstFragment", func(b *testing.B) {
|
b.Run("FirstFragment", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
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)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -545,7 +545,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("SecondFragment", func(b *testing.B) {
|
b.Run("SecondFragment", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
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)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -590,7 +590,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("200 HopByHop headers", func(b *testing.B) {
|
b.Run("200 HopByHop headers", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
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)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user