Compare commits

..
Author SHA1 Message Date
dependabot[bot] 49dc0cd34b Bump github.com/stretchr/testify from 1.12.0 to 1.12.1
Bumps [github.com/stretchr/testify](https://github.com/stretchr/testify) from 1.12.0 to 1.12.1.
- [Release notes](https://github.com/stretchr/testify/releases)
- [Commits](https://github.com/stretchr/testify/compare/v1.12.0...v1.12.1)

---
updated-dependencies:
- dependency-name: github.com/stretchr/testify
  dependency-version: 1.12.1
  dependency-type: direct:production
  update-type: version-update:semver-patch
...

Signed-off-by: dependabot[bot] <support@github.com>
2026-08-24 19:54:28 +00:00
99 changed files with 335 additions and 1267 deletions
-21
View File
@@ -41,24 +41,3 @@ jobs:
run: make fips140-all GOALS=smoke-docker
timeout-minutes: 10
smoke-self:
name: Run self traffic smoke test on macOS
runs-on: macos-latest
steps:
- uses: actions/checkout@v7
- uses: actions/setup-go@v7
with:
go-version: '1.26'
check-latest: true
- name: build
run: make bin
- name: run smoke-self
working-directory: ./.github/workflows/smoke
run: ./smoke-self.sh
timeout-minutes: 10
-130
View File
@@ -1,130 +0,0 @@
#!/bin/bash
# A host must be able to reach its own overlay address. Where the kernel sends
# that traffic through the tun rather than over loopback, nebula sees it and
# hands it straight back (immediatelyForwardToSelf), and whether the kernel
# accepts what comes back is only answerable against a real kernel. Runs one
# nebula on this machine as root and aims every probe at its own address.
set -e -x
set -o pipefail
V4=192.0.2.1
V6=2001:db8::1
case "$(uname -s)" in
Darwin) TUN_DEV=utun ;;
*) TUN_DEV=tun0 ;;
esac
ROOT="$(cd ../../.. && pwd)"
rm -rf build/self
mkdir -p build/self
cd build/self
cleanup() {
echo
echo " *** cleanup"
echo
set +e
if [ -n "$NEBULA_PID" ]
then
sudo kill "$NEBULA_PID"
fi
{ kill $(jobs -p); wait; } 2>/dev/null
sed 's/^/ [self] /' nebula.log
}
trap cleanup EXIT
# perl is on every platform this runs on; timeout(1) is not.
alarm() {
perl -e 'alarm shift; exec @ARGV' "$@"
}
RESULTS=""
FAILED=""
probe() {
local name="$1"
shift
if "$@"
then
RESULTS="$RESULTS $name=ok"
else
RESULTS="$RESULTS $name=FAIL"
FAILED="$FAILED $name"
fi
}
# Send one datagram, then wait for the listener to have written it out.
udp_probe() {
echo self | alarm 5 nc -u -w1 "$1" 3000 || true
set +x
for _ in $(seq 1 20)
do
if grep -q self "$2"
then
set -x
return 0
fi
sleep 0.25
done
set -x
return 1
}
"$ROOT/nebula-cert" ca -name "Smoke Test"
"$ROOT/nebula-cert" sign -name self -networks "$V4/24,$V6/64"
HOST=self AM_LIGHTHOUSE=true TUN_DEV="$TUN_DEV" ../../genconfig.sh >self.yml
"$ROOT/nebula" -config self.yml -test
sudo -v
sudo "$ROOT/nebula" -config self.yml >nebula.log 2>&1 &
NEBULA_PID=$!
for _ in $(seq 1 40)
do
ifconfig | grep "inet6 $V6 " >/dev/null && break
sleep 0.25
done
ifconfig | grep "inet $V4 "
ifconfig | grep "inet6 $V6 "
nc -l "$V4" 2000 >/dev/null &
nc -l "$V6" 2000 >/dev/null &
nc -u -l "$V4" 3000 >udp4.txt &
nc -u -l "$V6" 3000 >udp6.txt &
sleep 1
set +x
echo
echo " *** Testing self traffic from $V4"
echo
set -x
probe icmp4 alarm 5 ping -c1 "$V4"
probe tcp4 alarm 5 nc -z "$V4" 2000
probe udp4 udp_probe "$V4" udp4.txt
set +x
echo
echo " *** Testing self traffic from $V6"
echo
set -x
probe icmp6 alarm 5 ping6 -c1 "$V6"
probe tcp6 alarm 5 nc -z "$V6" 2000
probe udp6 udp_probe "$V6" udp6.txt
set +x
echo
echo " *** self traffic:$RESULTS"
echo
if [ -n "$FAILED" ]
then
echo "self traffic failed:$FAILED" >&2
exit 1
fi
+1 -4
View File
@@ -338,13 +338,10 @@ smoke-relay-docker: bin-docker
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
smoke-docker-ipv6: smoke-docker
smoke-self: bin
cd .github/workflows/smoke/ && ./smoke-self.sh
smoke-vagrant/%: bin-docker build/%/nebula
cd .github/workflows/smoke/ && ./build.sh $*
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
.FORCE:
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile debug docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 smoke-self test test-pkcs11 test-cov-html vet smoke-vagrant/%
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin bin-windows bin-windows-arm64 bin-darwin bin-freebsd bin-freebsd-arm64 bin-boringcrypto bin-fips140 bin-pkcs11 bin-docker boringcrypto build-test-mobile debug docker e2e e2ev e2evv e2evvv e2evvvv e2e-bench fips140 fips140-all $(ALL_GOFIPS140:%=fips140-%) install proto release release-linux release-freebsd release-openbsd release-netbsd release-boringcrypto release-fips140 service smoke-docker smoke-relay-docker smoke-docker-ipv6 test test-pkcs11 test-cov-html vet smoke-vagrant/%
.DEFAULT_GOAL := bin
+11 -2
View File
@@ -79,7 +79,13 @@ func (b *Bits) clearRange(startPos, count uint64) uint64 {
// handle the potential partial word before pos becomes u64 aligned
word := pos >> 6
bit := pos & 63
take := min(min(uint64(64)-bit, remaining), b.length-pos)
take := uint64(64) - bit
if take > remaining {
take = remaining
}
if take > b.length-pos {
take = b.length - pos
}
var mask uint64
if take == 64 {
mask = math.MaxUint64
@@ -183,7 +189,10 @@ func (b *Bits) Update(l *slog.Logger, i uint64) bool {
func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
// If i is a jump, adjust the window, record lost, update current, and return true
if i > b.current {
end := min(i, b.current+b.length)
end := i
if end > b.current+b.length {
end = b.current + b.length
}
count := end - b.current
startPos := (b.current + 1) & b.lengthMask
+1 -1
View File
@@ -120,7 +120,7 @@ func TestCertificate_SignP256_AlwaysNormalized(t *testing.T) {
pub := elliptic.Marshal(elliptic.P256(), priv.PublicKey.X, priv.PublicKey.Y)
rawPriv := priv.D.FillBytes(make([]byte, 32))
for i := range 1000 {
for i := 0; i < 1000; i++ {
if i&1 == 1 {
tbs.Version = Version1
} else {
+4 -4
View File
@@ -151,7 +151,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
var groups []string
if *cf.groups != "" {
for rg := range strings.SplitSeq(*cf.groups, ",") {
for _, rg := range strings.Split(*cf.groups, ",") {
g := strings.TrimSpace(rg)
if g != "" {
groups = append(groups, g)
@@ -171,7 +171,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
}
if *cf.networks != "" {
for rs := range strings.SplitSeq(*cf.networks, ",") {
for _, rs := range strings.Split(*cf.networks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -193,7 +193,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
}
if *cf.unsafeNetworks != "" {
for rs := range strings.SplitSeq(*cf.unsafeNetworks, ",") {
for _, rs := range strings.Split(*cf.unsafeNetworks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -221,7 +221,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
if !isP11 && *cf.encryption {
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 {
for range 5 {
for i := 0; i < 5; i++ {
errOut.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword()
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows
// +build !windows
package main
+4 -4
View File
@@ -146,7 +146,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 {
// ask for a passphrase until we get one
for range 5 {
for i := 0; i < 5; i++ {
errOut.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword()
@@ -203,7 +203,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
}
if *sf.networks != "" {
for rs := range strings.SplitSeq(*sf.networks, ",") {
for _, rs := range strings.Split(*sf.networks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -228,7 +228,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
}
if *sf.unsafeNetworks != "" {
for rs := range strings.SplitSeq(*sf.unsafeNetworks, ",") {
for _, rs := range strings.Split(*sf.unsafeNetworks, ",") {
rs := strings.Trim(rs, " ")
if rs != "" {
n, err := netip.ParsePrefix(rs)
@@ -247,7 +247,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
var groups []string
if *sf.groups != "" {
for rg := range strings.SplitSeq(*sf.groups, ",") {
for _, rg := range strings.Split(*sf.groups, ",") {
g := strings.TrimSpace(rg)
if g != "" {
groups = append(groups, g)
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows
// +build !windows
package main
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows
// +build !windows
package main
+1
View File
@@ -1,4 +1,5 @@
//go:build !linux
// +build !linux
package main
+9 -6
View File
@@ -5,7 +5,6 @@ import (
"errors"
"fmt"
"log/slog"
"maps"
"math"
"os"
"os/signal"
@@ -155,7 +154,9 @@ func (c *C) ReloadConfig() {
defer c.reloadLock.Unlock()
c.oldSettings = make(map[string]any)
maps.Copy(c.oldSettings, c.Settings)
for k, v := range c.Settings {
c.oldSettings[k] = v
}
err := c.Load(c.path)
if err != nil {
@@ -176,7 +177,9 @@ func (c *C) ReloadConfigString(raw string) error {
defer c.reloadLock.Unlock()
c.oldSettings = make(map[string]any)
maps.Copy(c.oldSettings, c.Settings)
for k, v := range c.Settings {
c.oldSettings[k] = v
}
err := c.LoadString(raw)
if err != nil {
@@ -213,7 +216,7 @@ func (c *C) GetStringSlice(k string, d []string) []string {
}
v := make([]string, len(rv))
for i := range v {
for i := 0; i < len(v); i++ {
v[i] = fmt.Sprintf("%v", rv[i])
}
@@ -307,8 +310,8 @@ func (c *C) IsSet(k string) bool {
}
func (c *C) get(k string, v any) any {
parts := strings.SplitSeq(k, ".")
for p := range parts {
parts := strings.Split(k, ".")
for _, p := range parts {
m, ok := v.(map[string]any)
if !ok {
return nil
+8 -14
View File
@@ -105,18 +105,11 @@ func (cm *connectionManager) getInactivityTimeout() time.Duration {
}
func (cm *connectionManager) In(h *HostInfo) {
h.markIn()
h.in.Store(true)
}
// OutNoRebind records outbound traffic without consuming the rebind epoch, for relayed sends: the direct path
// to the relay consumes the edge, the via send must not.
func (cm *connectionManager) OutNoRebind(h *HostInfo) {
h.markOutOnly()
}
// Out records outbound traffic and reports whether we rebound since this tunnel last sent
func (cm *connectionManager) Out(h *HostInfo) bool {
return h.markOut(cm.intf.rebindEpoch.Load())
func (cm *connectionManager) Out(h *HostInfo) {
h.out.Store(true)
}
func (cm *connectionManager) RelayUsed(localIndex uint32) {
@@ -135,7 +128,8 @@ func (cm *connectionManager) RelayUsed(localIndex uint32) {
// getAndResetTrafficCheck returns if there was any inbound or outbound traffic within the last tick and
// resets the state for this local index
func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) (bool, bool) {
in, out := h.takeTraffic()
in := h.in.Swap(false)
out := h.out.Swap(false)
if in || out {
h.lastUsed = now
}
@@ -352,7 +346,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
"tunnelCheck", m{"state": "alive", "method": "passive"},
)
}
hostinfo.setPendingDeletion(false)
hostinfo.pendingDeletion.Store(false)
if mainHostInfo {
decision = tryRehandshake
@@ -375,7 +369,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
return decision, hostinfo, primary
}
if hostinfo.isPendingDeletion() {
if hostinfo.pendingDeletion.Load() {
// We have already sent a test packet and nothing was returned, this hostinfo is dead
hostinfo.logger(cm.l).Info("Tunnel status",
"tunnelCheck", m{"state": "dead", "method": "active"},
@@ -426,7 +420,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
}
}
hostinfo.setPendingDeletion(true)
hostinfo.pendingDeletion.Store(true)
cm.trafficTimer.Add(hostinfo.localIndexId, cm.pendingDeletionInterval)
return decision, hostinfo, nil
}
+36 -36
View File
@@ -86,25 +86,25 @@ func Test_NewConnectionManagerTest(t *testing.T) {
// We saw traffic out to vpnIp
nc.Out(hostinfo)
nc.In(hostinfo)
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.pendingDeletion.Load())
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.True(t, hostinfo.sentSinceCheck())
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
assert.True(t, hostinfo.out.Load())
assert.True(t, hostinfo.in.Load())
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
// Do another traffic check tick, this host should be pending deletion now
nc.Out(hostinfo)
assert.True(t, hostinfo.sentSinceCheck())
assert.True(t, hostinfo.out.Load())
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.True(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.True(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
@@ -168,33 +168,33 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
// We saw traffic out to vpnIp
nc.Out(hostinfo)
nc.In(hostinfo)
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
assert.True(t, hostinfo.sentSinceCheck())
assert.False(t, hostinfo.isPendingDeletion())
assert.True(t, hostinfo.in.Load())
assert.True(t, hostinfo.out.Load())
assert.False(t, hostinfo.pendingDeletion.Load())
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
// Do a traffic check tick, should not be pending deletion but should not have any in/out packets recorded
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
// Do another traffic check tick, this host should be pending deletion now
nc.Out(hostinfo)
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.True(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.True(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
// We saw traffic, should no longer be pending deletion
nc.In(hostinfo)
nc.doTrafficCheck(hostinfo.localIndexId, p, nb, out, time.Now())
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
}
@@ -326,31 +326,31 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
// Do a traffic check tick, in and out should be cleared but should not be pending deletion
nc.Out(hostinfo)
nc.In(hostinfo)
assert.True(t, hostinfo.sentSinceCheck())
assert.True(t, (hostinfo.state.Load()&stateIn != 0))
assert.True(t, hostinfo.out.Load())
assert.True(t, hostinfo.in.Load())
now := time.Now()
decision, _, _ := nc.makeTrafficDecision(hostinfo.localIndexId, now)
assert.Equal(t, tryRehandshake, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*5))
assert.Equal(t, doNothing, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
// Do another traffic check tick, should still not be pending deletion
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Second*10))
assert.Equal(t, doNothing, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
@@ -358,9 +358,9 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
decision, _, _ = nc.makeTrafficDecision(hostinfo.localIndexId, now.Add(time.Minute*10))
assert.Equal(t, closeTunnel, decision)
assert.Equal(t, now, hostinfo.lastUsed)
assert.False(t, hostinfo.isPendingDeletion())
assert.False(t, hostinfo.sentSinceCheck())
assert.False(t, (hostinfo.state.Load()&stateIn != 0))
assert.False(t, hostinfo.pendingDeletion.Load())
assert.False(t, hostinfo.out.Load())
assert.False(t, hostinfo.in.Load())
assert.Contains(t, nc.hostMap.Indexes, hostinfo.localIndexId)
assert.Contains(t, nc.hostMap.Hosts, hostinfo.vpnAddrs[0])
}
+2 -29
View File
@@ -12,7 +12,6 @@ import (
"github.com/slackhq/nebula/handshake"
"github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -97,7 +96,7 @@ func TestConnectionState_NextMessageCounter(t *testing.T) {
assert.Equal(t, RejectAfterMessages, cs.messageCounter.Load())
// Continued send attempts stay refused and the counter never wraps
for range 10 {
for i := 0; i < 10; i++ {
_, ok = cs.NextMessageCounter()
assert.False(t, ok)
}
@@ -118,33 +117,7 @@ func TestSendNoMetricsDropsExhausted(t *testing.T) {
// The crossing send is refused: it records an exhaustion drop and never reaches connectionManager.Out.
assert.Equal(t, int64(1), f.messageMetrics.txExhausted.Count())
assert.False(t, hostinfo.sentSinceCheck())
}
// TestSendNoMetricsCloseTunnelKeepsRebindEpoch pins that a closing tunnel does not consume a rebind, a later
// packet on a re-established tunnel still needs that edge to trigger the far-side punch.
func TestSendNoMetricsCloseTunnelKeepsRebindEpoch(t *testing.T) {
initR, _ := runTestHandshake(t)
ci, err := newConnectionStateFromResult(initR)
require.NoError(t, err)
f := &Interface{
l: test.NewLogger(),
messageMetrics: &MessageMetrics{txExhausted: metrics.NewCounter()},
writers: []udp.Conn{udp.NoopConn{}},
connectionManager: &connectionManager{},
}
hostinfo := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.1")}, ConnectionState: ci}
// Tunnel is on epoch 0, then we rebind.
hostinfo.markOut(0)
f.rebindEpoch.Add(1)
remote := netip.MustParseAddrPort("10.0.0.2:4242")
f.sendNoMetrics(header.CloseTunnel, 0, ci, hostinfo, remote, []byte{}, make([]byte, 12), make([]byte, mtu), 0)
// markOut at the new epoch still reports the move, so the edge was preserved.
assert.True(t, hostinfo.markOut(1), "a CloseTunnel send must not consume the rebind epoch")
assert.False(t, hostinfo.out.Load())
}
func TestNewConnectionStateFromResult(t *testing.T) {
+1 -1
View File
@@ -212,7 +212,7 @@ func (c *Control) RebindUDPServer() {
c.f.lightHouse.SendUpdate()
// Let the main interface know that we rebound so that underlying tunnels know to trigger punches from their remotes
c.f.rebindEpoch.Add(1)
c.f.rebindCount++
}
// ListHostmapHosts returns details about the actual or pending (handshaking) hostmap by vpn ip
+1 -1
View File
@@ -247,7 +247,7 @@ func TestControl_ConcurrentStopAndStart(t *testing.T) {
c, _, _ := newReadyControl(t)
var wg sync.WaitGroup
for range 2 {
for i := 0; i < 2; i++ {
wg.Go(func() { c.Stop() })
}
wg.Go(func() { _ = c.Start() })
-10
View File
@@ -123,16 +123,6 @@ func (c *Control) SetLocalAddrsFn(fn func(*LocalAllowList) []netip.Addr) {
c.f.lightHouse.localAddrsFn = fn
}
// GetRebindEpochFor returns the rebind epoch a tunnel last sent under, so a test can tell whether a send
// consumed the epoch edge without having to infer it from lighthouse traffic.
func (c *Control) GetRebindEpochFor(vpnAddr netip.Addr) (uint32, bool) {
h := c.f.hostMap.QueryVpnAddr(vpnAddr)
if h == nil {
return 0, false
}
return h.state.Load() >> stateEpochShift, true
}
func (c *Control) KillPendingTunnel(vpnIp netip.Addr) bool {
hostinfo := c.f.handshakeManager.QueryVpnAddr(vpnIp)
if hostinfo == nil {
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1 -57
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
@@ -222,60 +223,3 @@ func TestRebindAdvertisesNewAddressAfterMove(t *testing.T) {
lhControl.Stop()
myControl.Stop()
}
// A relayed send records traffic but must not consume the rebind epoch. If it does, the next direct send to the
// relay host sees the epoch already current and never requeries, so the far side is never told to punch at our
// new address. This pins the SendVia call site, which the unit tests cannot reach.
func TestRebindRequeriesAfterRelayedSend(t *testing.T) {
t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
// No lighthouse on purpose: it would hand out a direct address for them and nothing would relay.
// Long connection manager timers so it never fires a direct test packet at the relay tunnel and bumps its
// epoch mid-test, which is the only other thing that touches that tunnel and would flake the assertion below.
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24",
m{"relay": m{"use_relays": true}, "timers": m{"connection_alive_interval": 3600, "pending_deletion_interval": 3600}})
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
r := router.NewR(t, myControl, relayControl, theirControl)
defer r.RenderFlow()
myControl.Start()
relayControl.Start()
theirControl.Start()
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("establish")))
r.RouteForAllUntilTxTun(theirControl)
r.RouteFor(time.Millisecond * 500)
hi := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false)
require.NotNil(t, hi, "expected a tunnel to them")
require.NotEmpty(t, hi.CurrentRelaysToMe, "them must be reachable only via the relay for this test to mean anything")
// sendNoMetrics only reaches SendVia when there is no direct remote, so pin that too. Without this the test
// keeps passing while quietly sending direct and never exercising the relay path.
require.False(t, hi.CurrentRemote.IsValid(), "them must have no direct remote, otherwise SendVia is never called")
before, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
require.True(t, ok, "expected a tunnel to the relay")
myControl.RebindUDPServer()
// Traffic to them goes through SendVia on the relay tunnel. That must record traffic without consuming the
// relay tunnel's own epoch edge, which belongs to the direct path.
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("relayed")))
r.RouteForAllUntilTxTun(theirControl)
after, ok := myControl.GetRebindEpochFor(relayVpnIpNet[0].Addr())
require.True(t, ok)
assert.Equal(t, before, after,
"a relayed send consumed the relay tunnel's rebind epoch, so the next direct send will not requery")
myControl.Stop()
relayControl.Stop()
theirControl.Stop()
}
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package router
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package router
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package e2e
+14 -15
View File
@@ -21,7 +21,6 @@ import (
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
)
type FirewallInterface interface {
@@ -263,11 +262,11 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
}
switch proto {
case iputil.IPProtocolTCP:
case firewall.ProtoTCP:
fp = ft.TCP
case iputil.IPProtocolUDP:
case firewall.ProtoUDP:
fp = ft.UDP
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
case firewall.ProtoICMP, firewall.ProtoICMPv6:
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
if startPort != firewall.PortAny {
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
@@ -365,13 +364,13 @@ func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw Fi
proto = firewall.ProtoAny
startPort, endPort, err = parsePort(sPort)
case "tcp":
proto = iputil.IPProtocolTCP
proto = firewall.ProtoTCP
startPort, endPort, err = parsePort(sPort)
case "udp":
proto = iputil.IPProtocolUDP
proto = firewall.ProtoUDP
startPort, endPort, err = parsePort(sPort)
case "icmp":
proto = iputil.IPProtocolICMP
proto = firewall.ProtoICMP
startPort = firewall.PortAny
endPort = firewall.PortAny
if sPort != "" {
@@ -561,9 +560,9 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
}
switch fp.Protocol {
case iputil.IPProtocolTCP:
case firewall.ProtoTCP:
c.Expires = time.Now().Add(f.TCPTimeout)
case iputil.IPProtocolUDP:
case firewall.ProtoUDP:
c.Expires = time.Now().Add(f.UDPTimeout)
default:
c.Expires = time.Now().Add(f.DefaultTimeout)
@@ -583,9 +582,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
c := &conn{}
switch fp.Protocol {
case iputil.IPProtocolTCP:
case firewall.ProtoTCP:
timeout = f.TCPTimeout
case iputil.IPProtocolUDP:
case firewall.ProtoUDP:
timeout = f.UDPTimeout
default:
timeout = f.DefaultTimeout
@@ -636,15 +635,15 @@ func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedC
}
switch p.Protocol {
case iputil.IPProtocolTCP:
case firewall.ProtoTCP:
if ft.TCP.match(p, incoming, c, caPool) {
return true
}
case iputil.IPProtocolUDP:
case firewall.ProtoUDP:
if ft.UDP.match(p, incoming, c, caPool) {
return true
}
case iputil.IPProtocolICMP, iputil.IPProtocolICMPv6:
case firewall.ProtoICMP, firewall.ProtoICMPv6:
if ft.ICMP.match(p, incoming, c, caPool) {
return true
}
@@ -681,7 +680,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
if p.Protocol == iputil.IPProtocolICMP || p.Protocol == iputil.IPProtocolICMPv6 {
if p.Protocol == firewall.ProtoICMP || p.Protocol == firewall.ProtoICMPv6 {
// port numbers are re-used for connection tracking of ICMP,
// but we don't want to actually filter on them.
return fp[firewall.PortAny].match(p, c, caPool)
+1 -1
View File
@@ -22,7 +22,7 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
l: l,
cache: make(ConntrackCache, cacheLen),
}
for i := range cacheLen {
for i := 0; i < cacheLen; i++ {
c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{}
}
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
+10 -7
View File
@@ -4,14 +4,17 @@ import (
"encoding/json"
"fmt"
"net/netip"
"github.com/slackhq/nebula/iputil"
)
type m = map[string]any
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`
PortFragment = -1 // Special value for matching `port: fragment`
)
@@ -42,13 +45,13 @@ func (fp *Packet) Copy() *Packet {
func (fp Packet) MarshalJSON() ([]byte, error) {
var proto string
switch fp.Protocol {
case iputil.IPProtocolTCP:
case ProtoTCP:
proto = "tcp"
case iputil.IPProtocolICMP:
case ProtoICMP:
proto = "icmp"
case iputil.IPProtocolICMPv6:
case ProtoICMPv6:
proto = "icmpv6"
case iputil.IPProtocolUDP:
case ProtoUDP:
proto = "udp"
default:
proto = fmt.Sprintf("unknown %v", fp.Protocol)
+33 -34
View File
@@ -13,7 +13,6 @@ import (
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -73,20 +72,20 @@ func TestFirewall_AddRule(t *testing.T) {
ti6, err := netip.ParsePrefix("fd12::34/128")
require.NoError(t, err)
require.NoError(t, fw.AddRule(true, iputil.IPProtocolTCP, 1, 1, []string{}, "", "", "", "", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoTCP, 1, 1, []string{}, "", "", "", "", ""))
// An empty rule is 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.Hosts)
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", ""))
assert.Nil(t, fw.InRules.UDP[1].Any.Any)
assert.Contains(t, fw.InRules.UDP[1].Any.Groups[0].Groups, "g1")
assert.Empty(t, fw.InRules.UDP[1].Any.Hosts)
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 1, 1, []string{}, "h1", "", "", "", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 1, 1, []string{}, "h1", "", "", "", ""))
//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.Empty(t, fw.InRules.ICMP[firewall.PortAny].Any.Groups)
@@ -117,11 +116,11 @@ func TestFirewall_AddRule(t *testing.T) {
assert.True(t, ok)
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "ca-name", ""))
assert.Contains(t, fw.InRules.UDP[1].CANames, "ca-name")
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
require.NoError(t, fw.AddRule(true, iputil.IPProtocolUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
require.NoError(t, fw.AddRule(true, firewall.ProtoUDP, 1, 1, []string{"g1"}, "", "", "", "", "ca-sha"))
assert.Contains(t, fw.InRules.UDP[1].CAShas, "ca-sha")
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c)
@@ -186,7 +185,7 @@ func TestFirewall_Drop(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10,
RemotePort: 90,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
@@ -264,7 +263,7 @@ func TestFirewall_DropV6(t *testing.T) {
RemoteAddr: netip.MustParseAddr("fd12::34"),
LocalPort: 10,
RemotePort: 90,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
@@ -351,7 +350,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
Certificate: &dummyCert{},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolUDP}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoUDP}, true, c, cp))
}
})
@@ -361,7 +360,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
Certificate: &dummyCert{},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 1}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 1}, true, c, cp))
}
})
@@ -371,7 +370,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
}
ip := netip.MustParsePrefix("9.254.254.254/32")
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
}
})
b.Run("pass proto, port, fail on local CIDRv6", func(b *testing.B) {
@@ -380,7 +379,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
}
ip := netip.MustParsePrefix("fd99::99/128")
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: ip.Addr()}, true, c, cp))
}
})
@@ -393,7 +392,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
}
})
b.Run("pass proto, port, any local CIDRv6, fail all group, name, and cidr", func(b *testing.B) {
@@ -405,7 +404,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
}
})
@@ -418,7 +417,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, 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) {
@@ -430,7 +429,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
assert.False(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
assert.False(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
}
})
@@ -442,7 +441,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}},
}
for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp))
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp))
}
})
@@ -454,7 +453,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}},
}
for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix.Addr()}, true, c, cp))
}
})
b.Run("pass on group on specific local cidr6", func(b *testing.B) {
@@ -465,7 +464,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"good-group": {}},
}
for n := 0; n < b.N; n++ {
assert.True(b, ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
assert.True(b, ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 100, LocalAddr: pfix6.Addr()}, true, c, cp))
}
})
@@ -477,7 +476,7 @@ func BenchmarkFirewallTable_match(b *testing.B) {
InvertedGroups: map[string]struct{}{"nope": {}},
}
for n := 0; n < b.N; n++ {
ft.match(firewall.Packet{Protocol: iputil.IPProtocolTCP, LocalPort: 10}, true, c, cp)
ft.match(firewall.Packet{Protocol: firewall.ProtoTCP, LocalPort: 10}, true, c, cp)
}
})
}
@@ -493,7 +492,7 @@ func TestFirewall_Drop2(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10,
RemotePort: 90,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
@@ -551,7 +550,7 @@ func TestFirewall_Drop3(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 1,
RemotePort: 1,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
@@ -639,7 +638,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
RemoteAddr: netip.MustParseAddr("fd12::34"),
LocalPort: 1,
RemotePort: 1,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
@@ -676,7 +675,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
LocalPort: 10,
RemotePort: 90,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
network := netip.MustParsePrefix("1.2.3.4/24")
@@ -759,13 +758,13 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
templ := firewall.Packet{
LocalAddr: netip.MustParseAddr("1.2.3.4"),
RemoteAddr: netip.MustParseAddr("1.2.3.4"),
Protocol: iputil.IPProtocolICMP,
Protocol: firewall.ProtoICMP,
Fragment: false,
}
t.Run("ICMP allowed", func(t *testing.T) {
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
require.NoError(t, fw.AddRule(true, iputil.IPProtocolICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoICMP, 0, 0, []string{"any"}, "", "", "", "", ""))
t.Run("zero ports", func(t *testing.T) {
p := templ.Copy()
p.LocalPort = 0
@@ -911,7 +910,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
RemoteAddr: netip.MustParseAddr("192.0.2.3"),
LocalPort: 1,
RemotePort: 1,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
@@ -962,7 +961,7 @@ func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) {
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
LocalPort: 443,
RemotePort: 55000,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
}
require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil),
@@ -1032,7 +1031,7 @@ func BenchmarkFirewallDropConntrackHit(b *testing.B) {
RemoteAddr: netip.MustParseAddr("192.0.2.2"),
LocalPort: 443,
RemotePort: 55000,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
}
cases := []struct {
@@ -1318,28 +1317,28 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
mf := &mockFirewall{}
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))
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding udp rule
conf = config.NewC(test.NewLogger())
mf = &mockFirewall{}
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))
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding icmp rule
conf = config.NewC(test.NewLogger())
mf = &mockFirewall{}
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))
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding icmp rule no port
conf = config.NewC(test.NewLogger())
mf = &mockFirewall{}
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
assert.Equal(t, addRuleCall{incoming: false, proto: iputil.IPProtocolICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
// Test adding any rule
conf = config.NewC(test.NewLogger())
@@ -1583,7 +1582,7 @@ func buildTestCase(setup testsetup, err error, theirPrefixes ...netip.Prefix) te
RemoteAddr: theirPrefixes[0].Addr(),
LocalPort: 10,
RemotePort: 90,
Protocol: iputil.IPProtocolUDP,
Protocol: firewall.ProtoUDP,
Fragment: false,
}
return testcase{
+1 -1
View File
@@ -19,7 +19,7 @@ require (
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
github.com/stretchr/testify v1.12.0
github.com/stretchr/testify v1.12.1
github.com/vishvananda/netlink v1.3.1
go.uber.org/goleak v1.3.0
go.yaml.in/yaml/v3 v3.0.5
+2 -2
View File
@@ -136,8 +136,8 @@ github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXf
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.12.0 h1:K6Mr6jO9JICuend/5xzTM03ydSV3vdNRYAdPSukj8uI=
github.com/stretchr/testify v1.12.0/go.mod h1:bOYBZb5qJ00vPzWfIqBUZPaxK8jWiXc6d3ErP4Ca9Gw=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
github.com/vishvananda/netlink v1.3.1/go.mod h1:ARtKouGSTGchR8aMwmkzC0qiNPrrWO5JS/XMVl45+b4=
github.com/vishvananda/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zdEY=
+13 -67
View File
@@ -239,15 +239,11 @@ const (
type HostInfo struct {
remote atomic.Pointer[netip.AddrPort]
remotes *RemoteList
promoteCounter atomic.Uint32
ConnectionState *ConnectionState
// Traffic bits, pendingDeletion, and the rebind epoch we last sent under
state atomic.Uint32
promoteCounter atomic.Uint32
remoteIndexId uint32
localIndexId uint32
remotes *RemoteList
remoteIndexId uint32
localIndexId uint32
// vpnAddrs is a list of vpn addresses assigned to this host that are within our own vpn networks
// The host may have other vpn addresses that are outside our
@@ -266,6 +262,11 @@ type HostInfo struct {
// This is used to limit lighthouse re-queries in chatty clients
nextLHQuery atomic.Int64
// lastRebindCount is the other side of Interface.rebindCount, if these values don't match then we need to ask LH
// for a punch from the remote end of this tunnel. The goal being to prime their conntrack for our traffic just like
// with a handshake
lastRebindCount int8
// lastHandshakeTime records the time the remote side told us about at the stage when the handshake was completed locally
// Stage 1 packet will contain it if I am a responder, stage 2 packet if I am an initiator
// This is used to avoid an attack where a handshake packet is replayed after some time
@@ -274,6 +275,9 @@ type HostInfo struct {
lastRoam time.Time
lastRoamRemote netip.AddrPort
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
in, out, pendingDeletion atomic.Bool
// lastUsed tracks the last time ConnectionManager checked the tunnel and it was in use.
// This value will be behind against actual tunnel utilization in the hot path.
// This should only be used by the ConnectionManagers ticker routine.
@@ -665,7 +669,7 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
hm.Indexes[hostinfo.localIndexId] = hostinfo
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
hostinfo.markOut(f.rebindEpoch.Load())
hostinfo.out.Store(true)
if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests
f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval)
}
@@ -766,64 +770,6 @@ func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interfac
}
}
// Bits within HostInfo.state, everything above stateEpochShift is the epoch
const (
stateIn uint32 = 1 << iota
stateOut
statePendingDeletion
stateFlags = stateIn | stateOut | statePendingDeletion
// The epoch is the top 29 bits, it would take 2^29 rebinds to wrap and we will never get there
stateEpochShift = 3
)
// markIn records inbound traffic
func (i *HostInfo) markIn() {
if i.state.Load()&stateIn == 0 {
i.state.Or(stateIn)
}
}
// markOut records a send and reports whether the epoch moved, meaning we want a punch from the far side
func (i *HostInfo) markOut(epoch uint32) bool {
e := epoch << stateEpochShift
for {
old := i.state.Load()
if old&stateOut != 0 && old&^stateFlags == e {
return false
}
if i.state.CompareAndSwap(old, old&stateFlags|stateOut|e) {
return old&^stateFlags != e
}
}
}
// markOutOnly records a send without consuming the rebind epoch, for paths that cannot act on a requery
func (i *HostInfo) markOutOnly() {
if i.state.Load()&stateOut == 0 {
i.state.Or(stateOut)
}
}
// takeTraffic clears both traffic bits, leaving the epoch alone, and reports what they were
func (i *HostInfo) takeTraffic() (in bool, out bool) {
old := i.state.And(^(stateIn | stateOut))
return old&stateIn != 0, old&stateOut != 0
}
func (i *HostInfo) setPendingDeletion(v bool) {
if v {
i.state.Or(statePendingDeletion)
} else {
i.state.And(^statePendingDeletion)
}
}
func (i *HostInfo) isPendingDeletion() bool {
return i.state.Load()&statePendingDeletion != 0
}
func (i *HostInfo) GetCert() *cert.CachedCertificate {
if i.ConnectionState != nil {
return i.ConnectionState.peerCert
-46
View File
@@ -401,49 +401,3 @@ func TestHostMap_RelayState(t *testing.T) {
assert.Equal(t, []netip.Addr{}, h1.relayState.relays)
}
// sentSinceCheck reports whether anything has been sent since the connection manager last looked. Test only:
// production reads the out bit through takeTraffic on the connection manager tick.
func (i *HostInfo) sentSinceCheck() bool {
return i.state.Load()&stateOut != 0
}
func TestHostInfo_markOut(t *testing.T) {
h := &HostInfo{}
h.markOut(5) // stamped when the tunnel was added
// A tunnel already on the current epoch has nothing to report, which is what keeps a fresh tunnel from
// requerying on its first packet
assert.False(t, h.markOut(5), "an unchanged epoch should not report a move")
assert.True(t, h.sentSinceCheck(), "the send is still recorded as traffic")
// A rebind is observed exactly once, so we requery once per rebind
assert.True(t, h.markOut(6), "a bumped epoch should report a move")
assert.False(t, h.markOut(6), "the epoch move should only be reported once")
// Traffic and pendingDeletion live in the same word and must survive an epoch change
h.setPendingDeletion(true)
h.markIn()
assert.True(t, h.markOut(7))
assert.True(t, h.isPendingDeletion(), "pendingDeletion must survive an epoch change")
in, out := h.takeTraffic()
assert.True(t, in, "inbound traffic must survive an epoch change")
assert.True(t, out)
// Clearing the traffic bits leaves the epoch alone, otherwise an idle tunnel would requery forever
assert.False(t, h.markOut(7), "takeTraffic must not disturb the epoch")
}
// A relayed send records traffic but must leave the rebind epoch for the direct path to consume, otherwise
// relaying to a host swallows the requery that gets the far side punching at our new address.
func TestHostInfo_markOutOnly(t *testing.T) {
h := &HostInfo{}
h.markOut(5)
h.markOutOnly()
assert.True(t, h.sentSinceCheck(), "a relayed send is still outbound traffic")
assert.False(t, h.markOut(5), "a relayed send must not disturb the epoch")
assert.True(t, h.markOut(6), "a relayed send must not consume the epoch edge")
assert.False(t, h.markOut(6))
}
+18 -14
View File
@@ -58,9 +58,6 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Parse
// kernel as one giant blob; segment first so the loopback
// path sees one IP datagram per Write.
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
// The kernel may have left the transport checksum for hardware
// offload to finish; nothing between here and the tun will.
iputil.SetTransportChecksum(seg)
_, werr := f.queues[q].Write(seg)
return werr
})
@@ -163,18 +160,21 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
// One traffic-out mark covers every segment of the superpacket; doing it
// per segment in sendInsideEncrypt paid an atomic store up to ~45 extra
// times per TSO packet, inside writeLock under boring crypto.
//
// We rebound since this tunnel last sent, ask the lighthouse to get the far side punching at us again
if f.connectionManager.Out(hostinfo) {
f.connectionManager.Out(hostinfo)
remote := hostinfo.GetRemote()
if hostinfo.lastRebindCount != f.rebindCount {
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
hostinfo.lastRebindCount = f.rebindCount
if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind epoch",
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
"vpnAddrs", hostinfo.vpnAddrs,
)
}
}
remote := hostinfo.GetRemote()
if !remote.IsValid() { //the relay path
//first, find our relay hostinfo:
var relayHostInfo *HostInfo
@@ -460,7 +460,7 @@ func (f *Interface) prepareSendVia(via *HostInfo,
}
out = header.Encode(out, header.Version, header.Message, header.MessageRelay, relay.RemoteIndex, c)
f.connectionManager.OutNoRebind(via)
f.connectionManager.Out(via)
// Authenticate the header and payload, but do not encrypt for this message type.
// The payload consists of the inner, unencrypted Nebula header, as well as the end-to-end encrypted payload.
@@ -553,13 +553,17 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
//l.WithField("trace", string(debug.Stack())).Error("out Header ", &Header{Version, t, st, 0, hostinfo.remoteIndexId, c}, p)
out = header.Encode(out, header.Version, t, st, hostinfo.remoteIndexId, c)
// A closing tunnel is torn down right after this, so skip the connection manager entirely: no point recording
// traffic or asking the lighthouse for a punch. Otherwise, if we rebound since this tunnel last sent, ask the
// lighthouse to get the far side punching at us again.
if t != header.CloseTunnel && f.connectionManager.Out(hostinfo) {
f.connectionManager.Out(hostinfo)
// Query our LH if we haven't since the last time we've been rebound, this will cause the remote to punch against
// all our addrs and enable a faster roaming.
if t != header.CloseTunnel && hostinfo.lastRebindCount != f.rebindCount {
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
hostinfo.lastRebindCount = f.rebindCount
if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("Lighthouse update triggered for punch due to rebind epoch",
f.l.Debug("Lighthouse update triggered for punch due to rebind counter",
"vpnAddrs", hostinfo.vpnAddrs,
)
}
-265
View File
@@ -1,265 +0,0 @@
package nebula
import (
"encoding/binary"
"io"
"net/netip"
"testing"
"github.com/gaissmai/bart"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
const (
ipv4HeaderLen = 20
ipv6HeaderLen = 40
)
// capturingTun is a tio.Queue that records what is written to it. A queue that
// discards writes is indistinguishable from a packet that was never forwarded.
type capturingTun struct {
writes [][]byte
}
func (c *capturingTun) Read() ([]tio.Packet, error) { return nil, io.EOF }
func (c *capturingTun) Close() error { return nil }
func (c *capturingTun) Write(b []byte) (int, error) {
c.writes = append(c.writes, append([]byte(nil), b...))
return len(b), nil
}
func newSelfForwardInterface(myAddrs ...netip.Addr) (*Interface, *capturingTun) {
vpnAddrs := &bart.Lite{}
for _, a := range myAddrs {
vpnAddrs.Insert(netip.PrefixFrom(a, a.BitLen()))
}
tun := &capturingTun{}
return &Interface{
l: test.NewLogger(),
myVpnAddrsTable: vpnAddrs,
myBroadcastAddrsTable: &bart.Lite{},
queues: []tio.Queue{tun},
}, tun
}
func consumeInside(f *Interface, packet []byte) {
f.consumeInsidePacket(tio.Packet{Bytes: packet}, &firewall.ParsedPacket{}, make([]byte, 12), nil, make([]byte, mtu), 0, nil)
}
// l4Proto describes one upper-layer header for these tests: its IP next-header
// value, where its checksum field sits within the header, and how to build a
// minimal instance of it.
type l4Proto struct {
name string
nextHdr uint8
cksumAt int
build func() []byte
}
var (
tcpSyn = l4Proto{"tcp", iputil.IPProtocolTCP, 16, func() []byte {
h := make([]byte, 20)
binary.BigEndian.PutUint16(h[0:2], 49152)
binary.BigEndian.PutUint16(h[2:4], 443)
binary.BigEndian.PutUint32(h[4:8], 0x11223344) // sequence
h[12] = 5 << 4 // data offset, no options
h[13] = 0x02 // SYN
binary.BigEndian.PutUint16(h[14:16], 65535) // window
return h
}}
udpDatagram = l4Proto{"udp", iputil.IPProtocolUDP, 6, func() []byte {
h := make([]byte, 8+4)
binary.BigEndian.PutUint16(h[0:2], 49152)
binary.BigEndian.PutUint16(h[2:4], 53)
binary.BigEndian.PutUint16(h[4:6], uint16(len(h)))
copy(h[8:], "ping")
return h
}}
icmpEcho = l4Proto{"icmp", iputil.IPProtocolICMP, 2, func() []byte { return echoRequest(8) }}
icmpv6Echo = l4Proto{"icmpv6", iputil.IPProtocolICMPv6, 2, func() []byte { return echoRequest(128) }}
)
// echoRequest builds an echo request body. The type differs between ICMP and
// ICMPv6, the rest of the header does not.
func echoRequest(typ uint8) []byte {
h := make([]byte, 8)
h[0] = typ
binary.BigEndian.PutUint16(h[4:6], 0xbeef) // identifier
binary.BigEndian.PutUint16(h[6:8], 1) // sequence
return h
}
func buildIPv6(src, dst netip.Addr, p l4Proto) []byte {
l4 := p.build()
pkt := make([]byte, ipv6HeaderLen+len(l4))
pkt[0] = 0x60
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(l4)))
pkt[6] = p.nextHdr
pkt[7] = 64
copy(pkt[8:24], src.AsSlice())
copy(pkt[24:40], dst.AsSlice())
copy(pkt[ipv6HeaderLen:], l4)
if l4 := pkt[ipv6HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
sum := ipv6PseudoheaderSum(src, dst, uint32(p.nextHdr), uint32(len(l4)))
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
}
return pkt
}
func buildIPv4(src, dst netip.Addr, p l4Proto) []byte {
l4 := p.build()
pkt := make([]byte, ipv4HeaderLen+len(l4))
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(len(pkt)))
pkt[8] = 64
pkt[9] = p.nextHdr
copy(pkt[12:16], src.AsSlice())
copy(pkt[16:20], dst.AsSlice())
copy(pkt[ipv4HeaderLen:], l4)
if l4 := pkt[ipv4HeaderLen:]; p.nextHdr == iputil.IPProtocolTCP || p.nextHdr == iputil.IPProtocolUDP {
sum := sumBytes(pkt[12:20], uint32(p.nextHdr)+uint32(len(l4)))
binary.BigEndian.PutUint16(l4[p.cksumAt:], ^fold(sumBytes(l4, sum)))
}
return pkt
}
// ipv6PseudoheaderSum is the RFC 2460 section 8.1 pseudo-header sum: source,
// destination, a 32 bit upper-layer packet length and a 32 bit zero-padded next
// header. Kept local to the test so these assertions do not check nebula's
// checksum code against itself.
func ipv6PseudoheaderSum(src, dst netip.Addr, nextHeader, length uint32) uint32 {
var csum uint32
s, d := src.AsSlice(), dst.AsSlice()
for i := 0; i < 16; i += 2 {
csum += uint32(s[i])<<8 | uint32(s[i+1])
csum += uint32(d[i])<<8 | uint32(d[i+1])
}
return csum + length + nextHeader
}
func sumBytes(b []byte, csum uint32) uint32 {
for i := 0; i+1 < len(b); i += 2 {
csum += uint32(b[i])<<8 | uint32(b[i+1])
}
if len(b)%2 == 1 {
csum += uint32(b[len(b)-1]) << 8
}
return csum
}
func fold(csum uint32) uint16 {
for csum > 0xffff {
csum = (csum >> 16) + (csum & 0xffff)
}
return uint16(csum)
}
// l4ChecksumValid6 verifies an IPv6 upper-layer checksum the way a receiver
// does: the pseudo-header plus the whole upper-layer segment, checksum field
// included, folds to 0xffff. The next header field is the upper-layer protocol
// only while there are no extension headers, which is all this file builds.
func l4ChecksumValid6(pkt []byte) bool {
src, _ := netip.AddrFromSlice(pkt[8:24])
dst, _ := netip.AddrFromSlice(pkt[24:40])
l4 := pkt[ipv6HeaderLen:]
return fold(sumBytes(l4, ipv6PseudoheaderSum(src, dst, uint32(pkt[6]), uint32(len(l4))))) == 0xffff
}
// l4ChecksumValid4 is the IPv4 counterpart: the RFC 793/768 pseudo-header is
// source, destination, a zero byte, the protocol and the upper-layer length.
func l4ChecksumValid4(pkt []byte) bool {
ihl := int(pkt[0]&0x0f) << 2
l4 := pkt[ihl:]
return fold(sumBytes(l4, sumBytes(pkt[12:20], uint32(pkt[9])+uint32(len(l4))))) == 0xffff
}
// TestConsumeInsidePacketSelfTraffic covers the self-addressed branch of
// consumeInsidePacket, taken where immediatelyForwardToSelf is set (see
// inside_bsd.go): the packet goes straight back to the tun, ahead of the
// firewall and the handshake.
func TestConsumeInsidePacketSelfTraffic(t *testing.T) {
v4 := netip.MustParseAddr("100.100.1.42")
v6 := netip.MustParseAddr("fd00::42")
tests := []struct {
name string
addr netip.Addr
pkt []byte
}{
{"ipv4/tcp", v4, buildIPv4(v4, v4, tcpSyn)},
{"ipv4/udp", v4, buildIPv4(v4, v4, udpDatagram)},
{"ipv4/icmp", v4, buildIPv4(v4, v4, icmpEcho)},
{"ipv6/tcp", v6, buildIPv6(v6, v6, tcpSyn)},
{"ipv6/udp", v6, buildIPv6(v6, v6, udpDatagram)},
{"ipv6/icmpv6", v6, buildIPv6(v6, v6, icmpv6Echo)},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
f, tun := newSelfForwardInterface(tt.addr)
// consumeInsidePacket writes through the slice it is handed, so a
// packet that arrived with a valid checksum must come back out of
// bytes taken before the call, unchanged.
want := append([]byte(nil), tt.pkt...)
consumeInside(f, tt.pkt)
if immediatelyForwardToSelf {
require.Len(t, tun.writes, 1)
assert.Equal(t, want, tun.writes[0])
} else {
assert.Empty(t, tun.writes, "self traffic reaches the tun over loopback here and must be dropped")
}
})
}
}
// TestConsumeInsidePacketSelfTrafficChecksum shows that the self-forward
// returns the bytes it was handed, so a packet that arrived with a wrong
// upper-layer checksum is written back with that same wrong checksum and the
// kernel drops it on re-entry.
//
// This is how a macOS host loses TCP and UDP to its own IPv6 overlay address:
// the kernel writes only the pseudo-header sum into the checksum field and
// defers completion to hardware offload, state that does not survive the
// crossing into userspace. Which kernels do this, for which protocols and IP
// versions, is a property of the kernel and belongs to a test against a live
// one; here the checksum is simply wrong, and the forward must make it right.
func TestConsumeInsidePacketSelfTrafficChecksum(t *testing.T) {
if !immediatelyForwardToSelf {
t.Skip("self traffic never reaches the tun on this platform")
}
versions := []struct {
name string
addr netip.Addr
build func(src, dst netip.Addr, p l4Proto) []byte
l4At int
valid func(pkt []byte) bool
}{
{"v4", netip.MustParseAddr("100.100.1.42"), buildIPv4, ipv4HeaderLen, l4ChecksumValid4},
{"v6", netip.MustParseAddr("fd00::42"), buildIPv6, ipv6HeaderLen, l4ChecksumValid6},
}
for _, v := range versions {
for _, p := range []l4Proto{tcpSyn, udpDatagram} {
t.Run(v.name+"/"+p.name, func(t *testing.T) {
pkt := v.build(v.addr, v.addr, p)
binary.BigEndian.PutUint16(pkt[v.l4At+p.cksumAt:], 0x1234)
require.False(t, v.valid(pkt), "the packet under test must start with a wrong checksum")
f, tun := newSelfForwardInterface(v.addr)
consumeInside(f, pkt)
require.Len(t, tun.writes, 1)
assert.True(t, v.valid(tun.writes[0]),
"a forwarded %s packet must carry a valid checksum, got 0x%04x",
p.name, binary.BigEndian.Uint16(tun.writes[0][v.l4At+p.cksumAt:]))
})
}
}
}
+2 -2
View File
@@ -107,8 +107,8 @@ type Interface struct {
sendRecvErrorConfig recvErrorConfig
acceptRecvErrorConfig recvErrorConfig
// Bumped on every udp rebind, tunnels compare it to decide they need a punch from the far side
rebindEpoch atomic.Uint32
// rebindCount is used to decide if an active tunnel should trigger a punch notification through a lighthouse
rebindCount int8
version string
conntrackCacheTimeout time.Duration
-146
View File
@@ -1,146 +0,0 @@
package iputil
import (
"encoding/binary"
"github.com/slackhq/nebula/overlay/checksum"
"golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
)
const udpHeaderLen = 8
// SetTransportChecksum recomputes the TCP or UDP checksum of an IPv4 or IPv6
// packet in place.
//
// A kernel that offloads checksums to the NIC hands a packet to a tun with the
// transport checksum unfinished: only the pseudo-header sum is in the field and
// the rest is left for hardware that a tun does not have. A packet written
// straight back to that tun is dropped on re-entry unless the checksum is
// completed first. ICMP is left alone; it arrived complete on the kernels this
// was measured against.
//
// So is any packet whose transport header cannot be located: fragments, unknown
// extension headers and truncated packets. An IPv6 fragment header is declined
// even when it carries the whole datagram (RFC 6946 atomic fragment), because
// the walk reports only that a fragment header was present.
func SetTransportChecksum(packet []byte) {
if len(packet) < 1 {
return
}
switch int(packet[0] >> 4) {
case ipv4.Version:
setTransportChecksum4(packet)
case ipv6.Version:
setTransportChecksum6(packet)
}
}
func setTransportChecksum4(packet []byte) {
if len(packet) < ipv4.HeaderLen {
return
}
ihl := int(packet[0]&0x0f) << 2
end := int(binary.BigEndian.Uint16(packet[2:4]))
if ihl < ipv4.HeaderLen || end < ihl || end > len(packet) {
return
}
// The checksum covers the whole datagram, which a fragment (MF set or a
// non-zero offset) does not carry.
if binary.BigEndian.Uint16(packet[6:8])&0x3fff != 0 {
return
}
transport, ok := transportExtent(packet[ihl:end], packet[9])
if !ok {
return
}
csum := ipv4PseudoheaderChecksum(packet[12:16], packet[16:20], uint32(packet[9]), uint32(len(transport)))
writeTransportChecksum(transport, packet[9], csum)
}
func setTransportChecksum6(packet []byte) {
if len(packet) < ipv6.HeaderLen {
return
}
end := ipv6.HeaderLen + int(binary.BigEndian.Uint16(packet[4:6]))
if end > len(packet) {
return
}
// The checksum covers the whole datagram, which a fragment does not carry.
// An unknown extension header hides where the transport header starts. A
// chain longer than the walk's budget ends it early, at an offset that was
// never checked against the packet.
proto, offset, _, anyFragment, err := IPv6FindUpperProtocol(packet[:end])
if err != nil || anyFragment || offset >= end {
return
}
transport, ok := transportExtent(packet[offset:end], proto)
if !ok {
return
}
csum := ipv6PseudoheaderChecksum(packet[8:24], packet[24:40], uint32(proto), uint32(len(transport)))
writeTransportChecksum(transport, proto, csum)
}
// transportExtent narrows a segment to the length its own header declares. UDP
// carries a Length field, and RFC 768 and RFC 8200 section 8.1 both make that
// field, not the IP payload extent, the length the pseudo-header counts and the
// checksum covers; a datagram padded out to a link's minimum frame is the usual
// way the two differ. TCP has no such field, so its segment runs to the end of
// the IP payload. A Length that overruns the bytes IP delivered describes a
// datagram that is not there.
func transportExtent(transport []byte, proto uint8) ([]byte, bool) {
if proto != IPProtocolUDP {
return transport, true
}
if len(transport) < udpHeaderLen {
return nil, false
}
ulen := int(binary.BigEndian.Uint16(transport[4:6]))
if ulen < udpHeaderLen || ulen > len(transport) {
return nil, false
}
return transport[:ulen], true
}
// writeTransportChecksum stores the checksum of transport, taken over the
// pseudo-header sum csum, in the header's checksum field. A UDP checksum that
// 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
// section 8.1).
func writeTransportChecksum(transport []byte, proto uint8, csum uint32) {
var at, minLen int
switch proto {
case IPProtocolTCP:
at, minLen = 16, 20
case IPProtocolUDP:
at, minLen = 6, udpHeaderLen
default:
return
}
if len(transport) < minLen {
return
}
transport[at], transport[at+1] = 0, 0
sum := ^checksum.Checksum(transport, fold(csum))
if sum == 0 && proto == IPProtocolUDP {
sum = 0xffff
}
binary.BigEndian.PutUint16(transport[at:], sum)
}
// fold reduces a pseudo-header sum to the 16 bit seed Checksum takes. Carrying
// the high half back into the low half is what keeps the reduction lossless, so
// the seed sums exactly as the wider value would; 0xffff is its fixed point.
// Every term of that sum comes from a 16 bit field, so it stays far below the
// width at which the accumulator would wrap.
func fold(csum uint32) uint16 {
for csum > 0xffff {
csum = (csum >> 16) + (csum & 0xffff)
}
return uint16(csum)
}
-242
View File
@@ -1,242 +0,0 @@
package iputil
import (
"encoding/binary"
"net"
"testing"
"github.com/google/gopacket"
"github.com/google/gopacket/layers"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/net/ipv6"
)
// serialize builds a packet with gopacket, whose checksums are computed
// independently of this package.
func serialize(t *testing.T, ls ...gopacket.SerializableLayer) []byte {
buf := gopacket.NewSerializeBuffer()
require.NoError(t, gopacket.SerializeLayers(buf, gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true}, ls...))
return append([]byte(nil), buf.Bytes()...)
}
// withExtensionHeader inserts an 8 byte IPv6 extension header of the given
// type between the IPv6 header and its payload. The transport checksum does not
// change: the pseudo-header counts only upper-layer bytes.
func withExtensionHeader(pkt []byte, typ layers.IPProtocol, hdr [8]byte) []byte {
hdr[0] = pkt[6]
out := make([]byte, 0, len(pkt)+8)
out = append(out, pkt[:40]...)
out = append(out, hdr[:]...)
out = append(out, pkt[40:]...)
out[6] = byte(typ)
binary.BigEndian.PutUint16(out[4:6], binary.BigEndian.Uint16(pkt[4:6])+8)
return out
}
// truncate copies the first n bytes into a buffer of exactly that capacity, so
// a read past the length panics instead of quietly succeeding.
func truncate(pkt []byte, n int) []byte {
out := make([]byte, n)
copy(out, pkt)
return out
}
// extChain builds an IPv6 packet fronted by n Destination Options headers. Each
// points at another one, so the walk spends its whole budget without reaching a
// transport header. lastExtLen inflates the final header's declared length,
// which is how the walk ends up past the end of the packet.
func extChain(n int, lastExtLen byte) []byte {
pkt := make([]byte, ipv6.HeaderLen)
pkt[0], pkt[6], pkt[7] = 0x60, 60, 64
for i := range n {
h := make([]byte, 8)
h[0] = 60
if i == n-1 {
h[1] = lastExtLen
}
pkt = append(pkt, h...)
}
pkt = append(pkt, make([]byte, 20)...)
binary.BigEndian.PutUint16(pkt[4:6], uint16(len(pkt)-ipv6.HeaderLen))
return pkt
}
func TestSetTransportChecksum(t *testing.T) {
// Source and destination differ so that a pseudo-header built from the wrong
// one, or from the two swapped, does not land on the same checksum anyway.
v4 := func(proto layers.IPProtocol) *layers.IPv4 {
return &layers.IPv4{Version: 4, TTL: 64, Id: 0x1234, Protocol: proto, SrcIP: net.IPv4(192, 0, 2, 1).To4(), DstIP: net.IPv4(198, 51, 100, 2).To4()}
}
v6 := func(proto layers.IPProtocol) *layers.IPv6 {
return &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: net.ParseIP("2001:db8::1"), DstIP: net.ParseIP("2001:db8:1::2")}
}
tcp := func(ip gopacket.NetworkLayer) *layers.TCP {
l := &layers.TCP{SrcPort: 49152, DstPort: 443, SYN: true, Window: 65535}
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
return l
}
udp := func(ip gopacket.NetworkLayer) *layers.UDP {
l := &layers.UDP{SrcPort: 49152, DstPort: 53}
require.NoError(t, l.SetNetworkLayerForChecksum(ip))
return l
}
payload := gopacket.Payload("self")
nop := layers.IPv4Option{OptionType: 1, OptionLength: 1}
ip4tcp := v4(layers.IPProtocolTCP)
ip4opts := v4(layers.IPProtocolTCP)
ip4opts.Options = []layers.IPv4Option{nop, nop, nop, nop}
ip4udp := v4(layers.IPProtocolUDP)
ip6tcp := v6(layers.IPProtocolTCP)
ip6udp := v6(layers.IPProtocolUDP)
hopByHop := [8]byte{0, 0, 1, 4} // next header, length 0, PadN of 4
// Bytes past the length the IP header declares are not part of the
// datagram and must not be summed.
trailing4 := append(serialize(t, ip4tcp, tcp(ip4tcp), payload), []byte("trailing")...)
trailing6 := append(serialize(t, ip6tcp, tcp(ip6tcp), payload), []byte("trailing")...)
// A datagram padded out past the length UDP declares: the pseudo-header
// counts the UDP Length field, so the checksum is the unpadded one.
padded4 := append(serialize(t, ip4udp, udp(ip4udp), payload), []byte("pad!")...)
binary.BigEndian.PutUint16(padded4[2:4], uint16(len(padded4)))
padded6 := append(serialize(t, ip6udp, udp(ip6udp), payload), []byte("pad!")...)
binary.BigEndian.PutUint16(padded6[4:6], uint16(len(padded6)-ipv6.HeaderLen))
// Corrupting the checksum and asking for it back must yield gopacket's
// packet, byte for byte.
recomputed := []struct {
name string
pkt []byte
cksum int
}{
{"v4 tcp", serialize(t, ip4tcp, tcp(ip4tcp), payload), 20 + 16},
{"v4 tcp with ip options", serialize(t, ip4opts, tcp(ip4opts), payload), 24 + 16},
{"v4 udp", serialize(t, ip4udp, udp(ip4udp), payload), 20 + 6},
{"v4 tcp header only", serialize(t, ip4tcp, tcp(ip4tcp)), 20 + 16},
{"v4 udp header only", serialize(t, ip4udp, udp(ip4udp)), 20 + 6},
{"v6 tcp", serialize(t, ip6tcp, tcp(ip6tcp), payload), 40 + 16},
{"v6 udp", serialize(t, ip6udp, udp(ip6udp), payload), 40 + 6},
{"v6 udp header only", serialize(t, ip6udp, udp(ip6udp)), 40 + 6},
{"v6 tcp behind hop-by-hop", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6HopByHop, hopByHop), 48 + 16},
{"v4 tcp with bytes past the total length", trailing4, 20 + 16},
{"v6 tcp with bytes past the payload length", trailing6, 40 + 16},
{"v4 udp padded past its declared length", padded4, 20 + 6},
{"v6 udp padded past its declared length", padded6, 40 + 6},
}
for _, tt := range recomputed {
t.Run(tt.name, func(t *testing.T) {
got := append([]byte(nil), tt.pkt...)
binary.BigEndian.PutUint16(got[tt.cksum:], 0x1234)
require.NotEqual(t, tt.pkt, got)
SetTransportChecksum(got)
assert.Equal(t, tt.pkt, got)
})
}
ip4frag := v4(layers.IPProtocolTCP)
ip4frag.Flags = layers.IPv4MoreFragments
ip4later := v4(layers.IPProtocolTCP)
ip4later.FragOffset = 1
ip4icmp := v4(layers.IPProtocolICMPv4)
badIHL := serialize(t, ip4tcp, tcp(ip4tcp), payload)
badIHL[0] = 0x44 // header length 16, shorter than an ipv4 header
shortTotalLen := serialize(t, ip4tcp, tcp(ip4tcp), payload)
binary.BigEndian.PutUint16(shortTotalLen[2:4], 10) // shorter than the header it introduces
cutTCP := serialize(t, ip4tcp, tcp(ip4tcp), payload)
binary.BigEndian.PutUint16(cutTCP[2:4], 20+19) // one byte short of a tcp header
cutTCP = truncate(cutTCP, 20+19)
cutUDP := serialize(t, ip4udp, udp(ip4udp), payload)
binary.BigEndian.PutUint16(cutUDP[2:4], 20+7) // one byte short of a udp header
cutUDP = truncate(cutUDP, 20+7)
// Two bytes short, so a transport header survives whole and the minimum
// length check cannot stand in for the bounds check.
cutV6 := truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 62)
fragment := [8]byte{0, 0, 0, 1, 0, 0, 0, 1} // next header, reserved, offset 0 with M set, id
overrun4 := serialize(t, ip4udp, udp(ip4udp), payload)
binary.BigEndian.PutUint16(overrun4[24:26], uint16(len(overrun4)-20+1)) // one byte past what ip delivered
overrun6 := serialize(t, ip6udp, udp(ip6udp), payload)
binary.BigEndian.PutUint16(overrun6[44:46], uint16(len(overrun6)-ipv6.HeaderLen+1))
shortUDPLen := serialize(t, ip4udp, udp(ip4udp), payload)
binary.BigEndian.PutUint16(shortUDPLen[24:26], 7) // shorter than the header it counts
// Where the checksum cannot be completed the packet is left as it came.
untouched := []struct {
name string
pkt []byte
cksum int
}{
{"v4 first fragment", serialize(t, ip4frag, tcp(ip4frag), payload), 20 + 16},
{"v4 later fragment", serialize(t, ip4later, tcp(ip4later), payload), 20 + 16},
{"v4 icmp", serialize(t, ip4icmp, &layers.ICMPv4{TypeCode: layers.CreateICMPv4TypeCode(8, 0), Id: 1, Seq: 1}, payload), 20 + 2},
{"v4 header length below the minimum", badIHL, 20 + 16},
{"v4 total length below the header length", shortTotalLen, 20 + 16},
{"v4 truncated below its total length", truncate(serialize(t, ip4tcp, tcp(ip4tcp), payload), 30), -1},
{"v4 tcp header cut short", cutTCP, 20 + 16},
{"v4 udp header cut short", cutUDP, -1},
{"v6 fragment", withExtensionHeader(serialize(t, ip6tcp, tcp(ip6tcp), payload), layers.IPProtocolIPv6Fragment, fragment), 48 + 16},
{"v6 truncated below its payload length", truncate(serialize(t, ip6tcp, tcp(ip6tcp), payload), 50), -1},
{"v6 truncated with a whole transport header still present", cutV6, 40 + 16},
{"v6 extension header chain longer than the walk", extChain(9, 0), 112 + 16},
{"v6 extension header chain running past the packet", extChain(8, 255), 104 + 16},
{"v4 udp length past the end of the datagram", overrun4, 20 + 6},
{"v6 udp length past the end of the datagram", overrun6, 40 + 6},
{"v4 udp length below a udp header", shortUDPLen, 20 + 6},
}
for _, tt := range untouched {
t.Run(tt.name, func(t *testing.T) {
if tt.cksum >= 0 {
binary.BigEndian.PutUint16(tt.pkt[tt.cksum:], 0x1234)
}
want := append([]byte(nil), tt.pkt...)
SetTransportChecksum(tt.pkt)
assert.Equal(t, want, tt.pkt)
})
}
t.Run("too short to carry a header", func(t *testing.T) {
for _, pkt := range [][]byte{nil, {}, {0x45}, {0x60}} {
assert.NotPanics(t, func() { SetTransportChecksum(pkt) })
}
})
t.Run("tcp checksum of zero goes out as zero", func(t *testing.T) {
pkt := serialize(t, ip4tcp, tcp(ip4tcp), gopacket.Payload{0, 0})
c := binary.BigEndian.Uint16(pkt[36:38])
require.NotZero(t, c)
// Only udp reserves zero to mean "not computed", so tcp keeps it.
binary.BigEndian.PutUint16(pkt[40:42], c)
SetTransportChecksum(pkt)
assert.Zero(t, binary.BigEndian.Uint16(pkt[36:38]))
})
t.Run("udp checksum of zero goes out as 0xffff", func(t *testing.T) {
pkt := serialize(t, ip4udp, udp(ip4udp), gopacket.Payload{0, 0})
c := binary.BigEndian.Uint16(pkt[26:28])
require.NotZero(t, c)
// The one's complement sum is now 0xffff - c; adding c to the payload
// makes it 0xffff, whose complement is zero.
binary.BigEndian.PutUint16(pkt[28:30], c)
SetTransportChecksum(pkt)
assert.Equal(t, uint16(0xffff), binary.BigEndian.Uint16(pkt[26:28]))
})
}
func TestFold(t *testing.T) {
// 0xffff is the fold's fixed point, so a loop bound one notch tight never
// terminates on it.
for _, tt := range []struct {
in uint32
want uint16
}{
{0, 0},
{0xffff, 0xffff},
{0x10000, 1},
{0x1fffe, 0xffff},
{0xffffffff, 0xffff},
} {
assert.Equal(t, tt.want, fold(tt.in))
}
}
-7
View File
@@ -27,13 +27,6 @@ const (
maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000
MaxRejectPacketSize = maxIPv6RejectPacketSize
IPProtocolICMP = 1
IPProtocolICMPv6 = 58
IPProtocolTCP = 6
IPProtocolUDP = 17
ICMPv6TypeEchoRequest = 128
ICMPv6TypeEchoReply = 129
)
func CreateRejectPacket(packet []byte, out []byte) []byte {
+8 -12
View File
@@ -244,12 +244,10 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
if pinThreads && routines > 1 && len(cpuAffinity) == 0 && !configTest {
// The operator didn't choose pin CPUs, so pick a default set that
// prefers performance cores and doesn't stack co-located instances
// onto allowed[0].
// key is used to seed the spreading of routines->cores.
// use PID if you want to ensure many different Nebulas in VMs or containers land on different cores
// use port if you want to always end up on the same cores, ideal for benchmarking.
key := uint64(os.Getpid()) //default to PID
// onto allowed[0]. The bound UDP port keys the per-instance spread:
// distinct across instances sharing a box, stable across restarts.
// A nil result keeps listenIn's stock allowed[i] fallback.
key := uint64(os.Getpid())
pinKeyStr := strings.ToLower(c.GetString("tun.pin_threads_key", ""))
switch pinKeyStr {
case "":
@@ -257,16 +255,14 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
case "pid":
l.Debug("tun.pin_threads_key is PID")
case "port":
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
l.Info("tun.pin_threads_key is port number")
key = uint64(ap.Port())
} else {
l.Warn("Failed to get a port number for tun.pin_threads_key, falling back to PID", "err", err)
}
l.Info("tun.pin_threads_key is port number")
default:
l.Warn("tun.pin_threads_key is invalid, using PID")
}
if ap, err := udpConns[0].LocalAddr(); err == nil && ap.Port() != 0 {
key = uint64(ap.Port())
}
cpuAffinity = cpupick.Default(routines, key, l)
}
+1
View File
@@ -1,4 +1,5 @@
//go:build boringcrypto
// +build boringcrypto
package noiseutil
+1
View File
@@ -1,4 +1,5 @@
//go:build boringcrypto
// +build boringcrypto
package noiseutil
+7 -6
View File
@@ -8,6 +8,7 @@ import (
"net/netip"
"time"
"github.com/google/gopacket/layers"
"golang.org/x/net/ipv6"
"github.com/slackhq/nebula/firewall"
@@ -368,15 +369,15 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
return nil
}
switch proto {
case iputil.IPProtocolICMPv6:
switch layers.IPProtocol(proto) {
case layers.IPProtocolICMPv6:
// An ICMPv6 message is at least type, code and checksum, 4 bytes. Only echo carries more than we read.
if dataLen < offset+4 {
return ErrIPv6PacketTooShort
}
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
switch data[offset] { //icmp type
case iputil.ICMPv6TypeEchoRequest, iputil.ICMPv6TypeEchoReply:
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply:
if dataLen < offset+6 {
return ErrIPv6PacketTooShort
}
@@ -385,7 +386,7 @@ func parseV6(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
fp.RemotePort = 0
}
case iputil.IPProtocolTCP, iputil.IPProtocolUDP:
case layers.IPProtocolTCP, layers.IPProtocolUDP:
if dataLen < offset+4 {
return ErrIPv6PacketTooShort
}
@@ -434,7 +435,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?
minLen := ihl
if !fp.Fragment {
if fp.Protocol == iputil.IPProtocolICMP {
if fp.Protocol == firewall.ProtoICMP {
minLen += minFwPacketLen + 2
} else {
minLen += minFwPacketLen
@@ -456,7 +457,7 @@ func parseV4(data []byte, incoming bool, fp *firewall.ParsedPacket) error {
if fp.Fragment {
fp.RemotePort = 0
fp.LocalPort = 0
} else if fp.Protocol == iputil.IPProtocolICMP { //note that orientation doesn't matter on ICMP
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
} else if incoming {
+19 -20
View File
@@ -9,7 +9,6 @@ import (
"github.com/google/gopacket"
"github.com/google/gopacket/layers"
"github.com/slackhq/nebula/iputil"
"github.com/slackhq/nebula/firewall"
"github.com/stretchr/testify/assert"
@@ -59,7 +58,7 @@ func Test_newPacket(t *testing.T) {
Src: net.IPv4(10, 0, 0, 1),
Dst: net.IPv4(10, 0, 0, 2),
Options: []byte{0, 1, 0, 2},
Protocol: iputil.IPProtocolTCP,
Protocol: firewall.ProtoTCP,
}
b, _ = h.Marshal()
@@ -67,7 +66,7 @@ func Test_newPacket(t *testing.T) {
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
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, uint16(3), p.RemotePort)
@@ -240,7 +239,7 @@ func Test_newPacket_v6(t *testing.T) {
// A good UDP packet
ip = layers.IPv6{
Version: 6,
NextHeader: iputil.IPProtocolUDP,
NextHeader: firewall.ProtoUDP,
HopLimit: 128,
SrcIP: net.IPv6linklocalallrouters,
DstIP: net.IPv6linklocalallnodes,
@@ -263,7 +262,7 @@ func Test_newPacket_v6(t *testing.T) {
// incoming
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(iputil.IPProtocolUDP), 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::1"), p.LocalAddr)
assert.Equal(t, uint16(36123), p.RemotePort)
@@ -273,7 +272,7 @@ func Test_newPacket_v6(t *testing.T) {
// outgoing
err = newPacket(b, false, p)
require.NoError(t, err)
assert.Equal(t, uint8(iputil.IPProtocolUDP), 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::1"), p.RemoteAddr)
assert.Equal(t, uint16(36123), p.LocalPort)
@@ -290,7 +289,7 @@ func Test_newPacket_v6(t *testing.T) {
// incoming
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(iputil.IPProtocolTCP), 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::1"), p.LocalAddr)
assert.Equal(t, uint16(36123), p.RemotePort)
@@ -300,7 +299,7 @@ func Test_newPacket_v6(t *testing.T) {
// outgoing
err = newPacket(b, false, p)
require.NoError(t, err)
assert.Equal(t, uint8(iputil.IPProtocolTCP), 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::1"), p.RemoteAddr)
assert.Equal(t, uint16(36123), p.LocalPort)
@@ -345,7 +344,7 @@ func Test_newPacket_v6(t *testing.T) {
err = newPacket(b, true, p)
require.NoError(t, err)
assert.Equal(t, uint8(iputil.IPProtocolUDP), 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::1"), p.LocalAddr)
assert.Equal(t, uint16(36123), p.RemotePort)
@@ -679,7 +678,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
pkt := make([]byte, realTCPAt+4)
pkt[0] = 0x60 // version 6
pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options
pkt[40] = byte(iputil.IPProtocolTCP) // Dest-Options NextHeader -> TCP
pkt[40] = byte(firewall.ProtoTCP) // Dest-Options NextHeader -> TCP
pkt[41] = 255 // HdrExtLen = 255
// Forged transport header at the pre-fix (wrong) offset: dst port 443.
@@ -688,7 +687,7 @@ func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) {
binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22)
require.NoError(t, newPacket(pkt, true, p))
assert.Equal(t, uint8(iputil.IPProtocolTCP), p.Protocol)
assert.Equal(t, uint8(firewall.ProtoTCP), p.Protocol)
// 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.
assert.Equal(t, uint16(22), p.LocalPort, "firewall must parse the real transport header, not the overflowed offset")
@@ -767,7 +766,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// Plain IPv4 TCP, IHL 20: L4 offset 20, no fragment shape.
v4 := make([]byte, 28)
v4[0] = 0x45
v4[9] = iputil.IPProtocolTCP
v4[9] = firewall.ProtoTCP
binary.BigEndian.PutUint16(v4[6:8], 0x4000) // DF only
require.NoError(t, newPacket(v4, true, p))
assert.Equal(t, 20, p.IPHdrLen)
@@ -778,7 +777,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// (Fragment false) but the coalescer must not touch it (FragAny true).
ff := make([]byte, 28)
ff[0] = 0x45
ff[9] = iputil.IPProtocolUDP
ff[9] = firewall.ProtoUDP
binary.BigEndian.PutUint16(ff[6:8], 0x2000) // MF, offset 0
require.NoError(t, newPacket(ff, true, p))
assert.False(t, p.Fragment)
@@ -788,7 +787,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// IPv4 non-first fragment (nonzero offset): both flags set.
nf := make([]byte, 28)
nf[0] = 0x45
nf[9] = iputil.IPProtocolUDP
nf[9] = firewall.ProtoUDP
binary.BigEndian.PutUint16(nf[6:8], 0x00b9)
require.NoError(t, newPacket(nf, true, p))
assert.True(t, p.Fragment)
@@ -797,7 +796,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// IPv4 with options (IHL 24): IPHdrLen tracks the real L4 offset.
opts := make([]byte, 32)
opts[0] = 0x46
opts[9] = iputil.IPProtocolTCP
opts[9] = firewall.ProtoTCP
binary.BigEndian.PutUint16(opts[6:8], 0x4000)
require.NoError(t, newPacket(opts, true, p))
assert.Equal(t, 24, p.IPHdrLen)
@@ -806,7 +805,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
// Plain IPv6 TCP: L4 at 40.
v6 := make([]byte, 60)
v6[0] = 0x60
v6[6] = iputil.IPProtocolTCP
v6[6] = firewall.ProtoTCP
require.NoError(t, newPacket(v6, true, p))
assert.Equal(t, 40, p.IPHdrLen)
assert.False(t, p.FragAny)
@@ -815,7 +814,7 @@ func Test_newPacket_parsedFields(t *testing.T) {
hbh := make([]byte, 60)
hbh[0] = 0x60
hbh[6] = 0 // hop-by-hop
hbh[40] = iputil.IPProtocolTCP
hbh[40] = firewall.ProtoTCP
hbh[41] = 0 // HdrExtLen 0 -> 8-byte header
require.NoError(t, newPacket(hbh, true, p))
assert.Equal(t, 48, p.IPHdrLen)
@@ -825,17 +824,17 @@ func Test_newPacket_parsedFields(t *testing.T) {
f6 := make([]byte, 60)
f6[0] = 0x60
f6[6] = 44 // fragment extension header
f6[40] = iputil.IPProtocolUDP
f6[40] = firewall.ProtoUDP
require.NoError(t, newPacket(f6, true, p))
assert.True(t, p.FragAny)
assert.False(t, p.Fragment)
assert.Equal(t, uint8(iputil.IPProtocolUDP), p.Protocol)
assert.Equal(t, uint8(firewall.ProtoUDP), p.Protocol)
// IPv6 non-first fragment: both set, walk stops at the fragment header.
f6n := make([]byte, 60)
f6n[0] = 0x60
f6n[6] = 44
f6n[40] = iputil.IPProtocolUDP
f6n[40] = firewall.ProtoUDP
binary.BigEndian.PutUint16(f6n[42:44], 0x0008)
require.NoError(t, newPacket(f6n, true, p))
assert.True(t, p.Fragment)
+2 -2
View File
@@ -127,7 +127,7 @@ func TestPseudoSumIPv6MatchesReference(t *testing.T) {
func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
rng := rand.New(rand.NewSource(0x1791))
for _, hdrLen := range []int{20, 24, 40, 60} {
for trial := range 200 {
for trial := 0; trial < 200; trial++ {
hdr := make([]byte, hdrLen)
rng.Read(hdr)
hdr[0] = 0x40 | byte(hdrLen/4)
@@ -156,7 +156,7 @@ func TestIPv4HdrChecksumMatchesReference(t *testing.T) {
// the way a receiver does (pseudo-header + L4 must sum to all-ones).
func TestChecksumSeedReceiverAcceptance(t *testing.T) {
rng := rand.New(rand.NewSource(0x1826))
for trial := range 200 {
for trial := 0; trial < 200; trial++ {
src := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
dst := [4]byte{byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256)), byte(rng.Intn(256))}
payLen := rng.Intn(1500)
+1 -1
View File
@@ -540,7 +540,7 @@ func TestCoalescerCapBySegments(t *testing.T) {
c := newTestTCPCoalescer(t, w)
pay := make([]byte, 512)
seq := uint32(1000)
for range tcpCoalesceMaxSegs + 5 {
for i := 0; i < tcpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildTCPv4(seq, tcpAck, pay)); err != nil {
t.Fatal(err)
}
+2 -2
View File
@@ -28,7 +28,7 @@ func TestSendBatchReserveCommitFlush(t *testing.T) {
b := NewSendBatch(fw, 4, 32)
ap := netip.MustParseAddrPort("10.0.0.1:4242")
for i := range 4 {
for i := 0; i < 4; i++ {
slot := b.Reserve(32)
if cap(slot) != 32 {
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
@@ -72,7 +72,7 @@ func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
b := NewSendBatch(fw, 3, 8)
ap := netip.MustParseAddrPort("10.0.0.1:80")
for i := range 3 {
for i := 0; i < 3; i++ {
s := b.Reserve(8)
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
b.Commit(pkt, ap)
+3 -3
View File
@@ -129,7 +129,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for range 3 {
for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
@@ -259,7 +259,7 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 100)
for range udpCoalesceMaxSegs + 5 {
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err)
}
@@ -317,7 +317,7 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
c := newTestUDPCoalescer(t, w)
pay := make([]byte, 1200)
for range 3 {
for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
t.Fatal(err)
}
+1 -1
View File
@@ -151,7 +151,7 @@ func TestChecksumTailPaths(t *testing.T) {
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
for k := 0; k <= maxK; k++ {
for tail := range 64 {
for tail := 0; tail < 64; tail++ {
length := 64*k + tail
for _, seed := range seeds {
for _, off := range offsets {
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package tio
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package tio
+7 -4
View File
@@ -1,4 +1,5 @@
//go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio
@@ -116,15 +117,17 @@ func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
}()
var wg sync.WaitGroup
for range writers {
wg.Go(func() {
for range perWriter {
for w := 0; w < writers; w++ {
wg.Add(1)
go func() {
defer wg.Done()
for i := 0; i < perWriter; i++ {
if _, werr := p.Write(payload); werr != nil {
t.Errorf("write: %v", werr)
return
}
}
})
}()
}
wg.Wait()
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package tio
+6 -5
View File
@@ -1,4 +1,5 @@
//go:build linux && !android && !e2e_testing
// +build linux,!android,!e2e_testing
package tio
@@ -136,7 +137,7 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
// payload
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
}
return pkt, virtio.NewHeader(
@@ -256,7 +257,7 @@ func TestSegmentTCPv6(t *testing.T) {
pkt[53] = 0x19 // FIN | ACK | PSH — exercise FIN clearing too
binary.BigEndian.PutUint16(pkt[54:56], 65535)
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i)
}
@@ -360,7 +361,7 @@ func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpLen+payLen)) // superpacket length
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff)
}
@@ -472,7 +473,7 @@ func TestSegmentUDPv6(t *testing.T) {
// Superpacket-wide length, as the kernel supplies it; see buildUSOv4.
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpLen+payLen))
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i)
}
@@ -771,7 +772,7 @@ func buildTSOv6(payLen, gso int) []byte {
pkt[53] = 0x10 // ACK only
binary.BigEndian.PutUint16(pkt[54:56], 65535)
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i)
}
return pkt
+1
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package virtio
+11 -4
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
// Package virtio implements the pure validation, header-correction, and
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
@@ -253,9 +254,12 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := range numSeg {
for i := 0; i < numSeg; i++ {
segStart := i * gsoSize
segEnd := min(segStart+gsoSize, payLen)
segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segPayLen := segEnd - segStart
segLen := headerLen + segPayLen
headerOff := i * gsoSize
@@ -355,9 +359,12 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := range numSeg {
for i := 0; i < numSeg; i++ {
segStart := i * gsoSize
segEnd := min(segStart+gsoSize, payLen)
segEnd := segStart + gsoSize
if segEnd > payLen {
segEnd = payLen
}
segPayLen := segEnd - segStart
segLen := headerLen + segPayLen
headerOff := i * gsoSize
+8 -7
View File
@@ -1,4 +1,5 @@
//go:build linux && !android
// +build linux,!android
package virtio
@@ -57,7 +58,7 @@ func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
pkt[33] = 0x18 // ACK | PSH
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + tcpLen, ipLen
@@ -81,7 +82,7 @@ func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
for i := range payLen {
for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + udpLen, ipLen
@@ -187,7 +188,7 @@ func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
// Payload bytes must be the original contiguous slice.
segPayLen := len(seg) - int(hdrLen)
wantPay := make([]byte, segPayLen)
for k := range segPayLen {
for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
@@ -316,7 +317,7 @@ func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
}
wantPay := make([]byte, segPayLen)
for k := range segPayLen {
for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
@@ -364,7 +365,7 @@ func buildUDPv4Single(payload []byte) (pkt []byte, hdr Hdr) {
// 0xffff, because all-zero is the reserved "no checksum" encoding that IPv6 rejects outright.
func TestFinishChecksumUDPZeroStoresAllOnes(t *testing.T) {
var payload []byte
for i := range 0x10000 {
for i := 0; i < 0x10000; i++ {
p := []byte{byte(i >> 8), byte(i)}
pkt, hdr := buildUDPv4Single(p)
cs, co := int(hdr.CsumStart), int(hdr.CsumOffset)
@@ -543,7 +544,7 @@ func TestBaseSumsMatchZeroingReference(t *testing.T) {
t.Run("ipv4", func(t *testing.T) {
for ihl := ipv4HeaderMinLen; ihl <= ipv4HeaderMaxLen; ihl += 4 {
for range 5000 {
for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, ihl)
for i := range pkt {
pkt[i] = randByte(&state)
@@ -573,7 +574,7 @@ func TestBaseSumsMatchZeroingReference(t *testing.T) {
for dataOff := 5; dataOff <= 15; dataOff++ {
tcpLen := dataOff * 4
headerLen := csumStart + tcpLen
for range 5000 {
for iter := 0; iter < 5000; iter++ {
pkt := make([]byte, headerLen+64)
for i := range pkt {
pkt[i] = randByte(&state)
+2 -2
View File
@@ -93,14 +93,14 @@ func prefixToMask(prefix netip.Prefix) netip.Addr {
}
func flipBytes(b []byte) []byte {
for i := range b {
for i := 0; i < len(b); i++ {
b[i] ^= 0xFF
}
return b
}
func orBytes(a []byte, b []byte) []byte {
ret := make([]byte, len(a))
for i := range a {
for i := 0; i < len(a); i++ {
ret[i] = a[i] | b[i]
}
return ret
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !ios && !e2e_testing
// +build !ios,!e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build ios && !e2e_testing
// +build ios,!e2e_testing
package overlay
+2 -1
View File
@@ -1,4 +1,5 @@
//go:build !android && !e2e_testing
// +build !android,!e2e_testing
package overlay
@@ -547,7 +548,7 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
if err != nil {
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
//retry twice more -- on some systems there appears to be a race condition where if we set routes too soon, netlink says `invalid argument`
for range 2 {
for i := 0; i < 2; i++ {
time.Sleep(100 * time.Millisecond)
err = netlink.RouteReplace(&nr)
if err == nil {
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !windows
// +build !windows
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package overlay
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package overlay
+2 -2
View File
@@ -108,7 +108,7 @@ func TestUserDeviceReadersConcurrentRace(t *testing.T) {
var wg sync.WaitGroup
run := func(idx int) {
defer wg.Done()
for range iterations {
for i := 0; i < iterations; i++ {
pkts, err := readers[idx].Read()
if err != nil {
errs <- err
@@ -136,7 +136,7 @@ func TestUserDeviceReadersConcurrentRace(t *testing.T) {
// waiting reader's private buffer, so reusing buf between writes is safe.
go func() {
buf := make([]byte, 32)
for i := range 2 * iterations {
for i := 0; i < 2*iterations; i++ {
for j := range buf {
buf[j] = byte(i + j)
}
+3 -3
View File
@@ -27,7 +27,7 @@ func TestPacketsAreBalancedEqually(t *testing.T) {
gw3count := 0
iterationCount := uint16(65535)
for i := range iterationCount {
for i := uint16(0); i < iterationCount; i++ {
packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"),
@@ -74,7 +74,7 @@ func TestPacketsAreBalancedByPriority(t *testing.T) {
gw2count := 0
iterationCount := uint16(65535)
for i := range iterationCount {
for i := uint16(0); i < iterationCount; i++ {
packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"),
@@ -115,7 +115,7 @@ func TestBalancePacketDistributsRandomlyAndReturnsFalseIfBucketsNotCalculated(t
gw1count := 0
gw2count := 0
for i := range iterationCount {
for i := uint16(0); i < iterationCount; i++ {
packet := firewall.Packet{
LocalAddr: netip.MustParseAddr("192.168.1.1"),
RemoteAddr: netip.MustParseAddr("10.0.0.1"),
+4 -5
View File
@@ -3,7 +3,6 @@ package routing
import (
"fmt"
"net/netip"
"strings"
)
const (
@@ -14,14 +13,14 @@ const (
type Gateways []Gateway
func (g Gateways) String() string {
var str strings.Builder
str := ""
for i, gw := range g {
str.WriteString(gw.String())
str += gw.String()
if i < len(g)-1 {
str.WriteString(", ")
str += ", "
}
}
return str.String()
return str
}
type Gateway struct {
+8 -4
View File
@@ -1,19 +1,21 @@
package nebula
import (
"context"
"testing"
"time"
)
func TestScheduler_PooledReuse(t *testing.T) {
ctx := t.Context()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
s := NewScheduler[int](16)
delivered := make(chan int, 256)
go s.Run(ctx, func(item int) { delivered <- item })
const N = 100
for i := range N {
for i := 0; i < N; i++ {
s.Schedule(ctx, i, time.Millisecond)
}
@@ -32,7 +34,8 @@ func TestScheduler_PooledReuse(t *testing.T) {
// BenchmarkScheduler_Schedule reports allocations per Schedule call.
// In steady state the Scheduler's sync.Pool means we should see zero allocs per op once the pool warms up.
func BenchmarkScheduler_Schedule(b *testing.B) {
ctx := b.Context()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
s := NewScheduler[int](b.N)
go s.Run(ctx, func(int) {})
@@ -48,7 +51,8 @@ func BenchmarkScheduler_Schedule(b *testing.B) {
// What we'd pay per Schedule if Punchy called time.AfterFunc directly without the pooled Scheduler.
// Allocates a *time.Timer plus a closure each call.
func BenchmarkBareAfterFunc(b *testing.B) {
ctx := b.Context()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
queue := make(chan int, b.N)
go func() {
+2 -2
View File
@@ -25,7 +25,7 @@ func AssertDeepCopyEqual(t *testing.T, a any, b any) {
}
func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name string) bool {
if v1.Type() == v2.Type() && v1.Type() == reflect.TypeFor[netip.Addr]() {
if v1.Type() == v2.Type() && v1.Type() == reflect.TypeOf(netip.Addr{}) {
// Ignore netip.Addr types since they reuse an interned global value
return false
}
@@ -72,7 +72,7 @@ func traverseDeepCopy(t *testing.T, v1 reflect.Value, v2 reflect.Value, name str
}
return traverseDeepCopy(t, v1.Elem(), v2.Elem(), name)
case reflect.Pointer:
case reflect.Ptr:
local := reflect.ValueOf(time.Local).Pointer()
if local == v1.Pointer() && local == v2.Pointer() {
return true
+1
View File
@@ -1,4 +1,5 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build darwin && !ios && !e2e_testing
// +build darwin,!ios,!e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !darwin || ios || e2e_testing
// +build !darwin ios e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
+5 -2
View File
@@ -294,7 +294,10 @@ func deliverSegments(r EncReader, from netip.AddrPort, payload []byte, segSize i
return
}
for off := 0; off < len(payload); off += segSize {
end := min(off+segSize, len(payload))
end := off + segSize
if end > len(payload) {
end = len(payload)
}
r(from, payload[off:end:end])
}
}
@@ -476,7 +479,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
return func() {
for i, gauges := range udpGauges {
if err := udpConns[i].(*StdConn).getMemInfo(&meminfo); err == nil {
for j := range unix.SK_MEMINFO_VARS {
for j := 0; j < unix.SK_MEMINFO_VARS; j++ {
gauges[j].Update(int64(meminfo[j]))
}
}
+4 -4
View File
@@ -107,7 +107,7 @@ func TestWriteBatchBadFamilyDeliversOthers(t *testing.T) {
got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 64)
for i := range 2 {
for i := 0; i < 2; i++ {
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
if rerr != nil {
t.Fatalf("expected 2 delivered packets, read #%d failed: %v", i+1, rerr)
@@ -161,7 +161,7 @@ func TestWriteBatchUnreachableDestDeliversOthers(t *testing.T) {
got := map[string]bool{}
rx.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 64)
for i := range 4 {
for i := 0; i < 4; i++ {
n, _, rerr := rx.ReadFromUDPAddrPort(buf)
if rerr != nil {
t.Fatalf("expected 4 delivered packets, read #%d failed: %v (got so far: %v)", i+1, rerr, got)
@@ -691,7 +691,7 @@ func TestGSOEngagesOnLoopback(t *testing.T) {
// The kernel must deliver the original datagram boundaries and bytes.
_ = rx.SetReadDeadline(time.Now().Add(5 * time.Second))
got := make([]byte, pktLen+1)
for i := range numPkts {
for i := 0; i < numPkts; i++ {
n, _, err := rx.ReadFromUDP(got)
if err != nil {
t.Fatalf("rx read %d: %v", i, err)
@@ -699,7 +699,7 @@ func TestGSOEngagesOnLoopback(t *testing.T) {
if n != pktLen {
t.Fatalf("rx read %d: len=%d want %d (kernel segmented at wrong boundary)", i, n, pktLen)
}
for j := range n {
for j := 0; j < n; j++ {
if got[j] != byte(i) {
t.Fatalf("rx read %d: byte %d = %#x, want %#x", i, j, got[j], byte(i))
}
+6 -3
View File
@@ -112,7 +112,7 @@ func (w *batchWriter) prepareWriteMessages(n int, offloadsEnabled bool) {
w.cmsg = make([]byte, n*w.cmsgSpace)
for k := range n {
for k := 0; k < n; k++ {
base := k * w.cmsgSpace
seg := (*unix.Cmsghdr)(unsafe.Pointer(&w.cmsg[base]))
seg.Level = unix.SOL_UDP
@@ -214,7 +214,7 @@ func (w *batchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort) (int, er
break
}
for k := range runLen {
for k := 0; k < runLen; k++ {
b := bufs[i+k]
if len(b) == 0 {
w.iovs[iovIdx+k].Base = nil
@@ -318,7 +318,10 @@ func (w *batchWriter) planRun(bufs [][]byte, addrs []netip.AddrPort, start, iovB
return 1, segSize
}
dst := addrs[start]
maxLen := min(iovBudget, w.maxGSOSegments)
maxLen := w.maxGSOSegments
if iovBudget < maxLen {
maxLen = iovBudget
}
runLen := 1
total := segSize
for runLen < maxLen && start+runLen < len(bufs) {
+2 -2
View File
@@ -64,13 +64,13 @@ func TestWriteBatchNoAllocs(t *testing.T) {
addrs = append(addrs, dst)
}
// GSO-eligible run with a short tail.
for range 8 {
for k := 0; k < 8; k++ {
add(payload, dstA)
}
add(short, dstA)
add(payload, dstA)
// Alternating destinations defeat coalescing entirely.
for k := range 4 {
for k := 0; k < 4; k++ {
dst := dstA
if k%2 == 0 {
dst = dstB
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
// Inspired by https://git.zx2c4.com/wireguard-go/tree/conn/bind_windows.go
+1
View File
@@ -1,4 +1,5 @@
//go:build e2e_testing
// +build e2e_testing
package udp
+1
View File
@@ -1,4 +1,5 @@
//go:build !e2e_testing
// +build !e2e_testing
package udp
+1 -1
View File
@@ -40,7 +40,7 @@ func AllowedCPUs() ([]int, error) {
return nil, err
}
cpus := make([]int, 0, set.Count())
for cpu := range len(set) * 64 {
for cpu := 0; cpu < len(set)*64; cpu++ {
if set.IsSet(cpu) {
cpus = append(cpus, cpu)
}