save about 10MB of RAM by not importing gopacket except for tests (#1864)
smoke-extra / freebsd-amd64 (push) Failing after 15s
smoke-extra / linux-amd64-ipv6disable (push) Failing after 15s
smoke-extra / netbsd-amd64 (push) Failing after 15s
smoke-extra / openbsd-amd64 (push) Failing after 16s
smoke-extra / linux-386 (push) Failing after 16s
smoke / Run multi node smoke test (push) Failing after 1m34s
Build and test / Static checks (push) Successful in 46s
Build and test / Test linux (push) Failing after 1m11s
Build and test / Test linux-pkcs11 (push) Failing after 1m51s
Build and test / Test linux-boringcrypto (push) Failing after 2m35s
Build and test / Test linux-fips140 (push) Failing after 2m41s
Build and test / Cross-build linux-arm (push) Successful in 2m57s
Build and test / Cross-build linux-mips (push) Successful in 4m0s
Build and test / Cross-build linux-other (push) Successful in 3m2s
Build and test / Cross-build windows (push) Successful in 1m2s
Build and test / Cross-build freebsd (push) Successful in 1m30s
Build and test / Cross-build netbsd (push) Successful in 1m29s
Build and test / Cross-build openbsd (push) Successful in 1m29s
Build and test / Cross-build mobile (push) Successful in 3m7s
smoke-extra / Run windows smoke test (push) Canceled after 0s
smoke / Run self traffic smoke test on macOS (push) Canceled after 0s
Build and test / Test macos (push) Canceled after 0s
Build and test / Test windows (push) Canceled after 0s
Build and test / CI status (push) Canceled after 0s

* save about 10MB of RAM by not importing gopacket except for tests

* big ol find-replace
This commit is contained in:
Jack Doan
2026-08-28 14:00:09 -05:00
committed by GitHub
parent ec3304e3a9
commit dd8f660c0a
8 changed files with 103 additions and 97 deletions
+15 -14
View File
@@ -21,6 +21,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/iputil"
) )
type FirewallInterface interface { type FirewallInterface interface {
@@ -262,11 +263,11 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
} }
switch proto { switch proto {
case firewall.ProtoTCP: case iputil.IPProtocolTCP:
fp = ft.TCP fp = ft.TCP
case firewall.ProtoUDP: case iputil.IPProtocolUDP:
fp = ft.UDP fp = ft.UDP
case firewall.ProtoICMP, firewall.ProtoICMPv6: case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided //ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
if startPort != firewall.PortAny { if startPort != firewall.PortAny {
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort) f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
@@ -364,13 +365,13 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
proto = firewall.ProtoAny proto = firewall.ProtoAny
startPort, endPort, err = parsePort(sPort) startPort, endPort, err = parsePort(sPort)
case "tcp": case "tcp":
proto = firewall.ProtoTCP proto = iputil.IPProtocolTCP
startPort, endPort, err = parsePort(sPort) startPort, endPort, err = parsePort(sPort)
case "udp": case "udp":
proto = firewall.ProtoUDP proto = iputil.IPProtocolUDP
startPort, endPort, err = parsePort(sPort) startPort, endPort, err = parsePort(sPort)
case "icmp": case "icmp":
proto = firewall.ProtoICMP proto = iputil.IPProtocolICMP
startPort = firewall.PortAny startPort = firewall.PortAny
endPort = firewall.PortAny endPort = firewall.PortAny
if sPort != "" { if sPort != "" {
@@ -560,9 +561,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
} }
switch fp.Protocol { switch fp.Protocol {
case firewall.ProtoTCP: case iputil.IPProtocolTCP:
c.Expires = time.Now().Add(f.TCPTimeout) c.Expires = time.Now().Add(f.TCPTimeout)
case firewall.ProtoUDP: case iputil.IPProtocolUDP:
c.Expires = time.Now().Add(f.UDPTimeout) c.Expires = time.Now().Add(f.UDPTimeout)
default: default:
c.Expires = time.Now().Add(f.DefaultTimeout) c.Expires = time.Now().Add(f.DefaultTimeout)
@@ -582,9 +583,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
c := &conn{} c := &conn{}
switch fp.Protocol { switch fp.Protocol {
case firewall.ProtoTCP: case iputil.IPProtocolTCP:
timeout = f.TCPTimeout timeout = f.TCPTimeout
case firewall.ProtoUDP: case iputil.IPProtocolUDP:
timeout = f.UDPTimeout timeout = f.UDPTimeout
default: default:
timeout = f.DefaultTimeout timeout = f.DefaultTimeout
@@ -635,15 +636,15 @@ func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedC
} }
switch p.Protocol { switch p.Protocol {
case firewall.ProtoTCP: case iputil.IPProtocolTCP:
if ft.TCP.match(p, incoming, c, caPool) { if ft.TCP.match(p, incoming, c, caPool) {
return true return true
} }
case firewall.ProtoUDP: case iputil.IPProtocolUDP:
if ft.UDP.match(p, incoming, c, caPool) { if ft.UDP.match(p, incoming, c, caPool) {
return true return true
} }
case firewall.ProtoICMP, firewall.ProtoICMPv6: case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
if ft.ICMP.match(p, incoming, c, caPool) { if ft.ICMP.match(p, incoming, c, caPool) {
return true return true
} }
@@ -680,7 +681,7 @@ func (fp firewallPort) match(p firewall.Packet, incoming bool, c *cert.CachedCer
} }
// this branch is here to catch traffic from FirewallTable.Any.match and FirewallTable.ICMP.match // this branch is here to catch traffic from FirewallTable.Any.match and FirewallTable.ICMP.match
if p.Protocol == firewall.ProtoICMP || p.Protocol == firewall.ProtoICMPv6 { if p.Protocol == iputil.IPProtocolICMP || p.Protocol == iputil.IPProtocolICMPv6 {
// port numbers are re-used for connection tracking of ICMP, // port numbers are re-used for connection tracking of ICMP,
// but we don't want to actually filter on them. // but we don't want to actually filter on them.
return fp[firewall.PortAny].match(p, c, caPool) return fp[firewall.PortAny].match(p, c, caPool)
+7 -10
View File
@@ -4,17 +4,14 @@ import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"net/netip" "net/netip"
"github.com/slackhq/nebula/iputil"
) )
type m = map[string]any type m = map[string]any
const ( const (
ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever ProtoAny = 0 // When we want to handle HOPOPT (0) we can change this, if ever
ProtoTCP = 6
ProtoUDP = 17
ProtoICMP = 1
ProtoICMPv6 = 58
PortAny = 0 // Special value for matching `port: any` PortAny = 0 // Special value for matching `port: any`
PortFragment = -1 // Special value for matching `port: fragment` PortFragment = -1 // Special value for matching `port: fragment`
) )
@@ -45,13 +42,13 @@ func (fp *Packet) Copy() *Packet {
func (fp Packet) MarshalJSON() ([]byte, error) { func (fp Packet) MarshalJSON() ([]byte, error) {
var proto string var proto string
switch fp.Protocol { switch fp.Protocol {
case ProtoTCP: case iputil.IPProtocolTCP:
proto = "tcp" proto = "tcp"
case ProtoICMP: case iputil.IPProtocolICMP:
proto = "icmp" proto = "icmp"
case ProtoICMPv6: case iputil.IPProtocolICMPv6:
proto = "icmpv6" proto = "icmpv6"
case ProtoUDP: case iputil.IPProtocolUDP:
proto = "udp" proto = "udp"
default: default:
proto = fmt.Sprintf("unknown %v", fp.Protocol) proto = fmt.Sprintf("unknown %v", fp.Protocol)
+34 -33
View File
@@ -13,6 +13,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/iputil"
"github.com/slackhq/nebula/test" "github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -72,20 +73,20 @@ func TestFirewall_AddRule(t *testing.T) {
ti6, err := netip.ParsePrefix("fd12::34/128") ti6, err := netip.ParsePrefix("fd12::34/128")
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, fw.AddRule(true, firewall.ProtoTCP, 1, 1, []string{}, "", "", "", "", "")) require.NoError(t, fw.AddRule(true, iputil.IPProtocolTCP, 1, 1, []string{}, "", "", "", "", ""))
// An empty rule is any // An empty rule is any
assert.True(t, fw.InRules.TCP[1].Any.Any.Any) assert.True(t, fw.InRules.TCP[1].Any.Any.Any)
assert.Empty(t, fw.InRules.TCP[1].Any.Groups) assert.Empty(t, fw.InRules.TCP[1].Any.Groups)
assert.Empty(t, fw.InRules.TCP[1].Any.Hosts) assert.Empty(t, fw.InRules.TCP[1].Any.Hosts)
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.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "")) require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
assert.Nil(t, fw.InRules.UDP[1].Any.Any) assert.Nil(t, fw.InRules.UDP[1].Any.Any)
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1") assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts) assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
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.ProtoICMP, 1, 1, []string{}, "h1", "", "", "", "")) require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 1, 1, []string{}, "h1", "", "", "", ""))
//no matter what port is given for icmp, it should end up as "any" //no matter what port is given for icmp, it should end up as "any"
assert.Nil(t, fw.InRules.ICMP[firewall.PortAny].Any.Any) assert.Nil(t, fw.InRules.ICMP[firewall.PortAny].Any.Any)
assert.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups) assert.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
@@ -116,11 +117,11 @@ func TestFirewall_AddRule(t *testing.T) {
assert.True(t, ok) assert.True(t, ok)
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.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", "")) require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name") assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
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.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha")) require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha") assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
@@ -185,7 +186,7 @@ func TestFirewall_Drop(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"), RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10, LocalPort: 10,
RemotePort: 90, RemotePort: 90,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
@@ -263,7 +264,7 @@ func TestFirewall_DropV6(t *testing.T) {
RemoteAddr: netip.MustParseAddr("fd12::34"), RemoteAddr: netip.MustParseAddr("fd12::34"),
LocalPort: 10, LocalPort: 10,
RemotePort: 90, RemotePort: 90,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
@@ -350,7 +351,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
Certificate: &dummyCert{}, Certificate: &dummyCert{},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoUDP}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolUDP}, true, c, cp))
} }
}) })
@@ -360,7 +361,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
Certificate: &dummyCert{}, Certificate: &dummyCert{},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 1}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 1}, true, c, cp))
} }
}) })
@@ -370,7 +371,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
} }
ip := netip.MustParsePrefix("9.254.254.254/32") ip := netip.MustParsePrefix("9.254.254.254/32")
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
} }
}) })
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) { b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
@@ -379,7 +380,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
} }
ip := netip.MustParsePrefix("fd99::99/128") ip := netip.MustParsePrefix("fd99::99/128")
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
} }
}) })
@@ -392,7 +393,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}}, InvertedGroups: map[string]struct{}{"nope": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
} }
}) })
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) { b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
@@ -404,7 +405,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}}, InvertedGroups: map[string]struct{}{"nope": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
} }
}) })
@@ -417,7 +418,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}}, InvertedGroups: map[string]struct{}{"nope": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
} }
}) })
b.Run("pass proto, port, specific local CIDRv6, fail all group, name, and cidr", func(b *testing.B) { b.Run("pass proto, port, specific local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
@@ -429,7 +430,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}}, InvertedGroups: map[string]struct{}{"nope": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp)) assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
} }
}) })
@@ -441,7 +442,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}}, InvertedGroups: map[string]struct{}{"good-group": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)) assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
} }
}) })
@@ -453,7 +454,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}}, InvertedGroups: map[string]struct{}{"good-group": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp)) assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
} }
}) })
b.Run("pass on group on specific local cidr6", func(b *testing.B) { b.Run("pass on group on specific local cidr6", func(b *testing.B) {
@@ -464,7 +465,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}}, InvertedGroups: map[string]struct{}{"good-group": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp)) assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
} }
}) })
@@ -476,7 +477,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}}, InvertedGroups: map[string]struct{}{"nope": {}},
} }
for n := 0; n < b.N; n++ { for n := 0; n < b.N; n++ {
ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp) ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp)
} }
}) })
} }
@@ -492,7 +493,7 @@ func TestFirewall_Drop2(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"), RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10, LocalPort: 10,
RemotePort: 90, RemotePort: 90,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
@@ -550,7 +551,7 @@ func TestFirewall_Drop3(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"), RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 1, LocalPort: 1,
RemotePort: 1, RemotePort: 1,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
@@ -638,7 +639,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
RemoteAddr: netip.MustParseAddr("fd12::34"), RemoteAddr: netip.MustParseAddr("fd12::34"),
LocalPort: 1, LocalPort: 1,
RemotePort: 1, RemotePort: 1,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
@@ -675,7 +676,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"), RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10, LocalPort: 10,
RemotePort: 90, RemotePort: 90,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
network := netip.MustParsePrefix("1.2.3.4/24") network := netip.MustParsePrefix("1.2.3.4/24")
@@ -758,13 +759,13 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
templ := firewall.Packet{ templ := firewall.Packet{
LocalAddr: netip.MustParseAddr("1.2.3.4"), LocalAddr: netip.MustParseAddr("1.2.3.4"),
RemoteAddr: netip.MustParseAddr("1.2.3.4"), RemoteAddr: netip.MustParseAddr("1.2.3.4"),
Protocol: firewall.ProtoICMP, Protocol: iputil.IPProtocolICMP,
Fragment: false, Fragment: false,
} }
t.Run("ICMP allowed", func(t *testing.T) { t.Run("ICMP allowed", func(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)
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 0, 0, []string{"any"}, "", "", "", "", "")) require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
t.Run("zero ports", func(t *testing.T) { t.Run("zero ports", func(t *testing.T) {
p := templ.Copy() p := templ.Copy()
p.LocalPort = 0 p.LocalPort = 0
@@ -910,7 +911,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
RemoteAddr: netip.MustParseAddr("192.0.2.3"), RemoteAddr: netip.MustParseAddr("192.0.2.3"),
LocalPort: 1, LocalPort: 1,
RemotePort: 1, RemotePort: 1,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP) assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
@@ -961,7 +962,7 @@ func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
RemoteAddr: netip.MustParseAddr("192.0.2.2"), RemoteAddr: netip.MustParseAddr("192.0.2.2"),
LocalPort: 443, LocalPort: 443,
RemotePort: 55000, RemotePort: 55000,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
} }
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil), require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
@@ -1031,7 +1032,7 @@ func BenchmarkFirewallDropConntrackHit(b *testing.B) {
RemoteAddr: netip.MustParseAddr("192.0.2.2"), RemoteAddr: netip.MustParseAddr("192.0.2.2"),
LocalPort: 443, LocalPort: 443,
RemotePort: 55000, RemotePort: 55000,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
} }
cases := []struct { cases := []struct {
@@ -1317,28 +1318,28 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
mf := &mockFirewall{} mf := &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}} conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf)) require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall) assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding udp rule // Test adding udp rule
conf = config.NewC(test.NewLogger()) conf = config.NewC(test.NewLogger())
mf = &mockFirewall{} mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}} conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf)) require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall) assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding icmp rule // Test adding icmp rule
conf = config.NewC(test.NewLogger()) conf = config.NewC(test.NewLogger())
mf = &mockFirewall{} mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}} conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf)) require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall) assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding icmp rule no port // Test adding icmp rule no port
conf = config.NewC(test.NewLogger()) conf = config.NewC(test.NewLogger())
mf = &mockFirewall{} mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}} conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf)) require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall) assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding any rule // Test adding any rule
conf = config.NewC(test.NewLogger()) conf = config.NewC(test.NewLogger())
@@ -1582,7 +1583,7 @@ func buildTestCase(setup testsetup, err error, theirPrefixes ...netip.Prefix) te
RemoteAddr: theirPrefixes[0].Addr(), RemoteAddr: theirPrefixes[0].Addr(),
LocalPort: 10, LocalPort: 10,
RemotePort: 90, RemotePort: 90,
Protocol: firewall.ProtoUDP, Protocol: iputil.IPProtocolUDP,
Fragment: false, Fragment: false,
} }
return testcase{ return testcase{
+7 -6
View File
@@ -8,6 +8,7 @@ import (
"github.com/gaissmai/bart" "github.com/gaissmai/bart"
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test" "github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
@@ -63,7 +64,7 @@ type l4Proto struct {
} }
var ( var (
tcpSyn = l4Proto{"tcp", firewall.ProtoTCP, 16, func() []byte { tcpSyn = l4Proto{"tcp", iputil.IPProtocolTCP, 16, func() []byte {
h := make([]byte, 20) h := make([]byte, 20)
binary.BigEndian.PutUint16(h[0:2], 49152) binary.BigEndian.PutUint16(h[0:2], 49152)
binary.BigEndian.PutUint16(h[2:4], 443) binary.BigEndian.PutUint16(h[2:4], 443)
@@ -74,7 +75,7 @@ var (
return h return h
}} }}
udpDatagram = l4Proto{"udp", firewall.ProtoUDP, 6, func() []byte { udpDatagram = l4Proto{"udp", iputil.IPProtocolUDP, 6, func() []byte {
h := make([]byte, 8+4) h := make([]byte, 8+4)
binary.BigEndian.PutUint16(h[0:2], 49152) binary.BigEndian.PutUint16(h[0:2], 49152)
binary.BigEndian.PutUint16(h[2:4], 53) binary.BigEndian.PutUint16(h[2:4], 53)
@@ -83,8 +84,8 @@ var (
return h return h
}} }}
icmpEcho = l4Proto{"icmp", firewall.ProtoICMP, 2, func() []byte { return echoRequest(8) }} icmpEcho = l4Proto{"icmp", iputil.IPProtocolICMP, 2, func() []byte { return echoRequest(8) }}
icmpv6Echo = l4Proto{"icmpv6", firewall.ProtoICMPv6, 2, func() []byte { return echoRequest(128) }} icmpv6Echo = l4Proto{"icmpv6", iputil.IPProtocolICMPv6, 2, func() []byte { return echoRequest(128) }}
) )
// echoRequest builds an echo request body. The type differs between ICMP and // echoRequest builds an echo request body. The type differs between ICMP and
@@ -107,7 +108,7 @@ func buildIPv6(src, dst netip.Addr, p l4Proto) []byte {
copy(pkt[8:24], src.AsSlice()) copy(pkt[8:24], src.AsSlice())
copy(pkt[24:40], dst.AsSlice()) copy(pkt[24:40], dst.AsSlice())
copy(pkt[ipv6HeaderLen:], l4) copy(pkt[ipv6HeaderLen:], l4)
if l4 := pkt[ipv6HeaderLen:]; p.nextHdr == firewall.ProtoTCP || p.nextHdr == firewall.ProtoUDP { if l4 := pkt[ipv6HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
sum := ipv6PseudoheaderSum(src, dst, uint32(p.nextHdr), uint32(len(l4))) sum := ipv6PseudoheaderSum(src, dst, uint32(p.nextHdr), uint32(len(l4)))
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum))) binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
} }
@@ -124,7 +125,7 @@ func buildIPv4(src, dst netip.Addr, p l4Proto) []byte {
copy(pkt[12:16], src.AsSlice()) copy(pkt[12:16], src.AsSlice())
copy(pkt[16:20], dst.AsSlice()) copy(pkt[16:20], dst.AsSlice())
copy(pkt[ipv4HeaderLen:], l4) copy(pkt[ipv4HeaderLen:], l4)
if l4 := pkt[ipv4HeaderLen:]; p.nextHdr == firewall.ProtoTCP || p.nextHdr == firewall.ProtoUDP { if l4 := pkt[ipv4HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
sum := sumBytes(pkt[12:20], uint32(p.nextHdr)+uint32(len(l4))) sum := sumBytes(pkt[12:20], uint32(p.nextHdr)+uint32(len(l4)))
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum))) binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
} }
+7 -8
View File
@@ -3,7 +3,6 @@ package iputil
import ( import (
"encoding/binary" "encoding/binary"
"github.com/google/gopacket/layers"
"github.com/slackhq/nebula/overlay/checksum" "github.com/slackhq/nebula/overlay/checksum"
"golang.org/x/net/ipv4" "golang.org/x/net/ipv4"
"golang.org/x/net/ipv6" "golang.org/x/net/ipv6"
@@ -57,7 +56,7 @@ func setTransportChecksum4(packet []byte) {
return return
} }
csum := ipv4PseudoheaderChecksum(packet[12:16], packet[16:20], uint32(packet[9]), uint32(len(transport))) csum := ipv4PseudoheaderChecksum(packet[12:16], packet[16:20], uint32(packet[9]), uint32(len(transport)))
writeTransportChecksum(transport, layers.IPProtocol(packet[9]), csum) writeTransportChecksum(transport, packet[9], csum)
} }
func setTransportChecksum6(packet []byte) { func setTransportChecksum6(packet []byte) {
@@ -83,7 +82,7 @@ func setTransportChecksum6(packet []byte) {
return return
} }
csum := ipv6PseudoheaderChecksum(packet[8:24], packet[24:40], uint32(proto), uint32(len(transport))) csum := ipv6PseudoheaderChecksum(packet[8:24], packet[24:40], uint32(proto), uint32(len(transport)))
writeTransportChecksum(transport, layers.IPProtocol(proto), csum) writeTransportChecksum(transport, proto, csum)
} }
// transportExtent narrows a segment to the length its own header declares. UDP // transportExtent narrows a segment to the length its own header declares. UDP
@@ -94,7 +93,7 @@ func setTransportChecksum6(packet []byte) {
// the IP payload. A Length that overruns the bytes IP delivered describes a // the IP payload. A Length that overruns the bytes IP delivered describes a
// datagram that is not there. // datagram that is not there.
func transportExtent(transport []byte, proto uint8) ([]byte, bool) { func transportExtent(transport []byte, proto uint8) ([]byte, bool) {
if layers.IPProtocol(proto) != layers.IPProtocolUDP { if proto != IPProtocolUDP {
return transport, true return transport, true
} }
if len(transport) < udpHeaderLen { if len(transport) < udpHeaderLen {
@@ -112,12 +111,12 @@ func transportExtent(transport []byte, proto uint8) ([]byte, bool) {
// computes to zero goes on the wire as 0xffff: zero means no checksum was // computes to zero goes on the wire as 0xffff: zero means no checksum was
// computed (RFC 768), and over IPv6 the checksum is mandatory (RFC 8200 // computed (RFC 768), and over IPv6 the checksum is mandatory (RFC 8200
// section 8.1). // section 8.1).
func writeTransportChecksum(transport []byte, proto layers.IPProtocol, csum uint32) { func writeTransportChecksum(transport []byte, proto uint8, csum uint32) {
var at, minLen int var at, minLen int
switch proto { switch proto {
case layers.IPProtocolTCP: case IPProtocolTCP:
at, minLen = 16, 20 at, minLen = 16, 20
case layers.IPProtocolUDP: case IPProtocolUDP:
at, minLen = 6, udpHeaderLen at, minLen = 6, udpHeaderLen
default: default:
return return
@@ -128,7 +127,7 @@ func writeTransportChecksum(transport []byte, proto layers.IPProtocol, csum uint
transport[at], transport[at+1] = 0, 0 transport[at], transport[at+1] = 0, 0
sum := ^checksum.Checksum(transport, fold(csum)) sum := ^checksum.Checksum(transport, fold(csum))
if sum == 0 && proto == layers.IPProtocolUDP { if sum == 0 && proto == IPProtocolUDP {
sum = 0xffff sum = 0xffff
} }
binary.BigEndian.PutUint16(transport[at:], sum) binary.BigEndian.PutUint16(transport[at:], sum)
+7
View File
@@ -27,6 +27,13 @@ const (
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000 maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
MaxRejectPacketSize = maxIPv6RejectPacketSize MaxRejectPacketSize = maxIPv6RejectPacketSize
IPProtocolICMP = 1
IPProtocolICMPv6 = 58
IPProtocolTCP = 6
IPProtocolUDP = 17
ICMPv6TypeEchoRequest = 128
ICMPv6TypeEchoReply = 129
) )
func CreateRejectPacket(packet []byte, out []byte) []byte { func CreateRejectPacket(packet []byte, out []byte) []byte {
+6 -7
View File
@@ -8,7 +8,6 @@ import (
"net/netip" "net/netip"
"time" "time"
"github.com/google/gopacket/layers"
"golang.org/x/net/ipv6" "golang.org/x/net/ipv6"
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
@@ -369,15 +368,15 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return nil return nil
} }
switch layers.IPProtocol(proto) { switch proto {
case layers.IPProtocolICMPv6: case iputil.IPProtocolICMPv6:
// An ICMPv6 message is at least type, code and checksum, 4 bytes. Only echo carries more than we read. // An ICMPv6 message is at least type, code and checksum, 4 bytes. Only echo carries more than we read.
if dataLen < offset+4 { if dataLen < offset+4 {
return ErrIPv6PacketTooShort return ErrIPv6PacketTooShort
} }
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6 fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
switch data[offset] { //icmp type switch data[offset] { //icmp type
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply: case iputil.ICMPv6TypeEchoRequest, iputil.ICMPv6TypeEchoReply:
if dataLen < offset+6 { if dataLen < offset+6 {
return ErrIPv6PacketTooShort return ErrIPv6PacketTooShort
} }
@@ -386,7 +385,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
fp.RemotePort = 0 fp.RemotePort = 0
} }
case layers.IPProtocolTCP, layers.IPProtocolUDP: case iputil.IPProtocolTCP, iputil.IPProtocolUDP:
if dataLen < offset+4 { if dataLen < offset+4 {
return ErrIPv6PacketTooShort return ErrIPv6PacketTooShort
} }
@@ -435,7 +434,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
// Accounting for a variable header length, do we have enough data for our src/dst tuples? // Accounting for a variable header length, do we have enough data for our src/dst tuples?
minLen := ihl minLen := ihl
if !fp.Fragment { if !fp.Fragment {
if fp.Protocol == firewall.ProtoICMP { if fp.Protocol == iputil.IPProtocolICMP {
minLen += minFwPacketLen + 2 minLen += minFwPacketLen + 2
} else { } else {
minLen += minFwPacketLen minLen += minFwPacketLen
@@ -457,7 +456,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
if fp.Fragment { if fp.Fragment {
fp.RemotePort = 0 fp.RemotePort = 0
fp.LocalPort = 0 fp.LocalPort = 0
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP } else if fp.Protocol == iputil.IPProtocolICMP { //note that orientation doesn't matter on ICMP
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
fp.LocalPort = 0 //code would be uint16(data[ihl+1]) fp.LocalPort = 0 //code would be uint16(data[ihl+1])
} else if incoming { } else if incoming {
+20 -19
View File
@@ -9,6 +9,7 @@ import (
"github.com/google/gopacket" "github.com/google/gopacket"
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
@@ -58,7 +59,7 @@ func Test_newPacket(t *testing.T) {
Src: net.IPv4(10, 0, 0, 1), Src: net.IPv4(10, 0, 0, 1),
Dst: net.IPv4(10, 0, 0, 2), Dst: net.IPv4(10, 0, 0, 2),
Options: []byte{0, 1, 0, 2}, Options: []byte{0, 1, 0, 2},
Protocol: firewall.ProtoTCP, Protocol: iputil.IPProtocolTCP,
} }
b, _ = h.Marshal() b, _ = h.Marshal()
@@ -66,7 +67,7 @@ func Test_newPacket(t *testing.T) {
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
assert.Equal(t, netip.MustParseAddr("10.0.0.2"), p.LocalAddr) assert.Equal(t, netip.MustParseAddr("10.0.0.2"), p.LocalAddr)
assert.Equal(t, netip.MustParseAddr("10.0.0.1"), p.RemoteAddr) assert.Equal(t, netip.MustParseAddr("10.0.0.1"), p.RemoteAddr)
assert.Equal(t, uint16(3), p.RemotePort) assert.Equal(t, uint16(3), p.RemotePort)
@@ -239,7 +240,7 @@ func Test_newPacket_v6(t *testing.T) {
// A good UDP packet // A good UDP packet
ip = layers.IPv6{ ip = layers.IPv6{
Version: 6, Version: 6,
NextHeader: firewall.ProtoUDP, NextHeader: iputil.IPProtocolUDP,
HopLimit: 128, HopLimit: 128,
SrcIP: net.IPv6linklocalallrouters, SrcIP: net.IPv6linklocalallrouters,
DstIP: net.IPv6linklocalallnodes, DstIP: net.IPv6linklocalallnodes,
@@ -262,7 +263,7 @@ func Test_newPacket_v6(t *testing.T) {
// incoming // incoming
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolUDP), 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)
assert.Equal(t, uint16(36123), p.RemotePort) assert.Equal(t, uint16(36123), p.RemotePort)
@@ -272,7 +273,7 @@ func Test_newPacket_v6(t *testing.T) {
// outgoing // outgoing
err = newPacket(b, false, p) err = newPacket(b, false, p)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
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)
assert.Equal(t, uint16(36123), p.LocalPort) assert.Equal(t, uint16(36123), p.LocalPort)
@@ -289,7 +290,7 @@ func Test_newPacket_v6(t *testing.T) {
// incoming // incoming
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolTCP), 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)
assert.Equal(t, uint16(36123), p.RemotePort) assert.Equal(t, uint16(36123), p.RemotePort)
@@ -299,7 +300,7 @@ func Test_newPacket_v6(t *testing.T) {
// outgoing // outgoing
err = newPacket(b, false, p) err = newPacket(b, false, p)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
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)
assert.Equal(t, uint16(36123), p.LocalPort) assert.Equal(t, uint16(36123), p.LocalPort)
@@ -344,7 +345,7 @@ func Test_newPacket_v6(t *testing.T) {
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolUDP), 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)
assert.Equal(t, uint16(36123), p.RemotePort) assert.Equal(t, uint16(36123), p.RemotePort)
@@ -678,7 +679,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
pkt := make([]byte, realTCPAt+4) pkt := make([]byte, realTCPAt+4)
pkt[0] = 0x60 // version 6 pkt[0] = 0x60 // version 6
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP pkt[40] = byte(iputil.IPProtocolTCP) // Dest-Options NextHeader -> TCP
pkt[41] = 255 // HdrExtLen = 255 pkt[41] = 255 // HdrExtLen = 255
// Forged transport header at the pre-fix (wrong) offset: dst port 443. // Forged transport header at the pre-fix (wrong) offset: dst port 443.
@@ -687,7 +688,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22) binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
require.NoError(t, newPacket(pkt, true, p)) require.NoError(t, newPacket(pkt, true, p))
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
// LocalPort is the destination port for incoming traffic. It must be the real port (22) // LocalPort is the destination port for incoming traffic. It must be the real port (22)
// the host delivers to, not the forged 443 at the overflowed offset. // the host delivers to, not the forged 443 at the overflowed offset.
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset") assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
@@ -766,7 +767,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape. // Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
v4 := make([]byte, 28) v4 := make([]byte, 28)
v4[0] = 0x45 v4[0] = 0x45
v4[9] = firewall.ProtoTCP v4[9] = iputil.IPProtocolTCP
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
require.NoError(t, newPacket(v4, true, p)) require.NoError(t, newPacket(v4, true, p))
assert.Equal(t, 20, p.IPHdrLen) assert.Equal(t, 20, p.IPHdrLen)
@@ -777,7 +778,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// (Fragment false) but the coalescer must not touch it (FragAny true). // (Fragment false) but the coalescer must not touch it (FragAny true).
ff := make([]byte, 28) ff := make([]byte, 28)
ff[0] = 0x45 ff[0] = 0x45
ff[9] = firewall.ProtoUDP ff[9] = iputil.IPProtocolUDP
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0 binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
require.NoError(t, newPacket(ff, true, p)) require.NoError(t, newPacket(ff, true, p))
assert.False(t, p.Fragment) assert.False(t, p.Fragment)
@@ -787,7 +788,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// IPv4 non-first fragment (nonzero offset): both flags set. // IPv4 non-first fragment (nonzero offset): both flags set.
nf := make([]byte, 28) nf := make([]byte, 28)
nf[0] = 0x45 nf[0] = 0x45
nf[9] = firewall.ProtoUDP nf[9] = iputil.IPProtocolUDP
binary.BigEndian.PutUint16(nf[6:8], 0x00b9) binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
require.NoError(t, newPacket(nf, true, p)) require.NoError(t, newPacket(nf, true, p))
assert.True(t, p.Fragment) assert.True(t, p.Fragment)
@@ -796,7 +797,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset. // IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
opts := make([]byte, 32) opts := make([]byte, 32)
opts[0] = 0x46 opts[0] = 0x46
opts[9] = firewall.ProtoTCP opts[9] = iputil.IPProtocolTCP
binary.BigEndian.PutUint16(opts[6:8], 0x4000) binary.BigEndian.PutUint16(opts[6:8], 0x4000)
require.NoError(t, newPacket(opts, true, p)) require.NoError(t, newPacket(opts, true, p))
assert.Equal(t, 24, p.IPHdrLen) assert.Equal(t, 24, p.IPHdrLen)
@@ -805,7 +806,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// Plain IPv6 TCP: L4 at 40. // Plain IPv6 TCP: L4 at 40.
v6 := make([]byte, 60) v6 := make([]byte, 60)
v6[0] = 0x60 v6[0] = 0x60
v6[6] = firewall.ProtoTCP v6[6] = iputil.IPProtocolTCP
require.NoError(t, newPacket(v6, true, p)) require.NoError(t, newPacket(v6, true, p))
assert.Equal(t, 40, p.IPHdrLen) assert.Equal(t, 40, p.IPHdrLen)
assert.False(t, p.FragAny) assert.False(t, p.FragAny)
@@ -814,7 +815,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
hbh := make([]byte, 60) hbh := make([]byte, 60)
hbh[0] = 0x60 hbh[0] = 0x60
hbh[6] = 0 // hop-by-hop hbh[6] = 0 // hop-by-hop
hbh[40] = firewall.ProtoTCP hbh[40] = iputil.IPProtocolTCP
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
require.NoError(t, newPacket(hbh, true, p)) require.NoError(t, newPacket(hbh, true, p))
assert.Equal(t, 48, p.IPHdrLen) assert.Equal(t, 48, p.IPHdrLen)
@@ -824,17 +825,17 @@ func Test_newPacket_parsedFields(t *testing.T) {
f6 := make([]byte, 60) f6 := make([]byte, 60)
f6[0] = 0x60 f6[0] = 0x60
f6[6] = 44 // fragment extension header f6[6] = 44 // fragment extension header
f6[40] = firewall.ProtoUDP f6[40] = iputil.IPProtocolUDP
require.NoError(t, newPacket(f6, true, p)) require.NoError(t, newPacket(f6, true, p))
assert.True(t, p.FragAny) assert.True(t, p.FragAny)
assert.False(t, p.Fragment) assert.False(t, p.Fragment)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol) assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
// IPv6 non-first fragment: both set, walk stops at the fragment header. // IPv6 non-first fragment: both set, walk stops at the fragment header.
f6n := make([]byte, 60) f6n := make([]byte, 60)
f6n[0] = 0x60 f6n[0] = 0x60
f6n[6] = 44 f6n[6] = 44
f6n[40] = firewall.ProtoUDP f6n[40] = iputil.IPProtocolUDP
binary.BigEndian.PutUint16(f6n[42:44], 0x0008) binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
require.NoError(t, newPacket(f6n, true, p)) require.NoError(t, newPacket(f6n, true, p))
assert.True(t, p.Fragment) assert.True(t, p.Fragment)