From 184cdc8586431a6e773bf37b7e71eb3d416f9b51 Mon Sep 17 00:00:00 2001 From: John Maguire Date: Tue, 23 Jun 2026 14:21:28 -0400 Subject: [PATCH 01/25] Add OCI image labels with version info (#1772) --- .github/workflows/release.yml | 5 ++++- docker/Dockerfile | 10 ++++++++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index e4ca2933..39ddfc5c 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -163,7 +163,10 @@ jobs: mkdir -p build/linux-{amd64,arm64} tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/ tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/ - docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 --tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}" + docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,linux/arm64 \ + --build-arg VERSION="${GITHUB_REF#refs/tags/v}" \ + --build-arg REVISION="${GITHUB_SHA}" \ + --tag "${DOCKER_IMAGE_REPO}:${DOCKER_IMAGE_TAG}" --tag "${DOCKER_IMAGE_REPO}:${GITHUB_REF#refs/tags/v}" release: name: Create and Upload Release diff --git a/docker/Dockerfile b/docker/Dockerfile index 400e275b..d705fce3 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -1,6 +1,16 @@ FROM gcr.io/distroless/static:latest ARG TARGETOS TARGETARCH + +ARG VERSION=dev +ARG REVISION=unknown +LABEL org.opencontainers.image.title="nebula" \ + org.opencontainers.image.description="A scalable overlay networking tool with a focus on performance, simplicity and security" \ + org.opencontainers.image.vendor="Nebula OSS" \ + org.opencontainers.image.source="https://github.com/slackhq/nebula" \ + org.opencontainers.image.version="${VERSION}" \ + org.opencontainers.image.revision="${REVISION}" + COPY build/$TARGETOS-$TARGETARCH/nebula /nebula COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert From 58ab7250f54f7d825062a0016d9b406bb94a9f27 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Mon, 29 Jun 2026 12:47:13 -0500 Subject: [PATCH 02/25] Bump actions/checkout from 6 to 7 (#1771) Bumps [actions/checkout](https://github.com/actions/checkout) from 6 to 7. - [Release notes](https://github.com/actions/checkout/releases) - [Changelog](https://github.com/actions/checkout/blob/main/CHANGELOG.md) - [Commits](https://github.com/actions/checkout/compare/v6...v7) --- updated-dependencies: - dependency-name: actions/checkout dependency-version: '7' dependency-type: direct:production update-type: version-update:semver-major ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- .github/workflows/release.yml | 10 +++++----- .github/workflows/smoke-extra.yml | 6 +++--- .github/workflows/smoke.yml | 2 +- .github/workflows/test.yml | 6 +++--- 4 files changed, 12 insertions(+), 12 deletions(-) diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 39ddfc5c..9c5b1c3e 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -10,7 +10,7 @@ jobs: name: Build Linux/BSD All runs-on: ubuntu-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -36,7 +36,7 @@ jobs: id-token: write contents: read steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -76,7 +76,7 @@ jobs: HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }} runs-on: macos-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -134,7 +134,7 @@ jobs: # be overwritten - name: Checkout code if: ${{ env.HAS_DOCKER_CREDS == 'true' }} - uses: actions/checkout@v6 + uses: actions/checkout@v7 - name: Download artifacts if: ${{ env.HAS_DOCKER_CREDS == 'true' }} @@ -173,7 +173,7 @@ jobs: needs: [build-linux, build-darwin, build-windows] runs-on: ubuntu-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - name: Download artifacts uses: actions/download-artifact@v8 diff --git a/.github/workflows/smoke-extra.yml b/.github/workflows/smoke-extra.yml index e0428e9c..8f71ead5 100644 --- a/.github/workflows/smoke-extra.yml +++ b/.github/workflows/smoke-extra.yml @@ -30,7 +30,7 @@ jobs: VAGRANT_DEFAULT_PROVIDER: libvirt steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -62,7 +62,7 @@ jobs: VAGRANT_DEFAULT_PROVIDER: virtualbox steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -88,7 +88,7 @@ jobs: runs-on: windows-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: diff --git a/.github/workflows/smoke.yml b/.github/workflows/smoke.yml index b3c5f2d9..9a4b0013 100644 --- a/.github/workflows/smoke.yml +++ b/.github/workflows/smoke.yml @@ -18,7 +18,7 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 2abb3740..269f0edb 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -18,7 +18,7 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -78,7 +78,7 @@ jobs: e2e-cmd: make e2evv steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -123,7 +123,7 @@ jobs: - {name: mobile, make-target: build-test-mobile} steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: From 02471b412112818ecdb2ed0d62188245d78c7b7f Mon Sep 17 00:00:00 2001 From: zorvios <35658698+zorvios@users.noreply.github.com> Date: Wed, 1 Jul 2026 22:12:14 +0100 Subject: [PATCH 03/25] wireshark: fix Lua 5.4 bitwise operation (#1776) The Nebula Wireshark dissector uses bit32.band() when parsing the Nebula packet type. On Wireshark builds using Lua 5.4, bit32 is not available, which causes dissection to fail with: Lua Error: nebula.lua:65: attempt to index a nil value (global 'bit32') Wireshark bundles Lua BitOp and exposes it globally as bit for Lua dissectors. This is the documented API for maximum backwards compatibility across supported Lua versions. Use bit.band() instead of bit32.band(). --- dist/wireshark/nebula.lua | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dist/wireshark/nebula.lua b/dist/wireshark/nebula.lua index d17dc7a0..5c7c17f1 100644 --- a/dist/wireshark/nebula.lua +++ b/dist/wireshark/nebula.lua @@ -62,7 +62,7 @@ function nebula.dissector(tvbuf, pktinfo, root) tree:add(pf_version, tvbuf:range(0,1)) local type = tree:add(pf_type, tvbuf:range(0,1)) - local nebula_type = bit32.band(tvbuf:range(0,1):uint(), 0x0F) + local nebula_type = bit.band(tvbuf:range(0,1):uint(), 0x0F) if nebula_type == 0 then local stage = tvbuf(8,8):uint64() tree:add(pf_subtype_handshake, tvbuf:range(1,1)) From 6afca0f46129dfb5102e438ada6e54467c2e022e Mon Sep 17 00:00:00 2001 From: Jack Doan Date: Fri, 3 Jul 2026 11:06:22 -0500 Subject: [PATCH 04/25] correctly handle a test packet with a payload longer than the header (#1778) --- e2e/echo_test.go | 85 ++++++++++++++++++++++++++++++++++++++++++++++++ outside.go | 3 +- 2 files changed, 87 insertions(+), 1 deletion(-) create mode 100644 e2e/echo_test.go diff --git a/e2e/echo_test.go b/e2e/echo_test.go new file mode 100644 index 00000000..5e1299b2 --- /dev/null +++ b/e2e/echo_test.go @@ -0,0 +1,85 @@ +//go:build e2e_testing +// +build e2e_testing + +package e2e + +import ( + "testing" + "time" + + "github.com/slackhq/nebula" + "github.com/slackhq/nebula/cert" + "github.com/slackhq/nebula/cert_test" + "github.com/slackhq/nebula/e2e/router" + "github.com/slackhq/nebula/header" + "github.com/slackhq/nebula/udp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func assertTestRequestEchoed(t *testing.T, cipher string) { + ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) + over := m{"cipher": cipher} + a, aNet, aUdp, _ := newSimpleServer(cert.Version1, ca, caKey, "a", "10.128.0.1/24", over) + b, bNet, bUdp, _ := newSimpleServer(cert.Version1, ca, caKey, "b", "10.128.0.2/24", over) + + a.InjectLightHouseAddr(bNet[0].Addr(), bUdp) + b.InjectLightHouseAddr(aNet[0].Addr(), aUdp) + a.Start() + b.Start() + t.Cleanup(func() { a.Stop(); b.Stop() }) + r := router.NewR(t, a, b) + defer r.RenderFlow() + + assertTunnel(t, aNet[0].Addr(), bNet[0].Addr(), a, b, r) + drainUDPTx(a) + drainUDPTx(b) + + payload := []byte("a test payload well over sixteen bytes long, wow it's so very long long long!") + require.Greater(t, len(payload), header.Len) + a.GetF().SendMessageToVpnAddr(header.Test, header.TestRequest, bNet[0].Addr(), payload, make([]byte, 12, 12), make([]byte, udp.MTU)) + + // Deliver A's request to B; B must echo a reply back + b.InjectUDPPacket(a.GetFromUDP(true)) + reply := nextUDPTxOfType(t, b, header.Test, header.TestReply, 2*time.Second) + + assert.Equal(t, aUdp, reply.To, "the reply must go back to the requester") + // header + echoed payload + 16-byte AEAD tag: proves the whole payload + // round-tripped rather than being dropped or truncated. + assert.Equal(t, header.Len+len(payload)+16, len(reply.Data), "the full payload must be echoed back") +} + +func TestTestRequestEchoesLongPayloadAES(t *testing.T) { + assertTestRequestEchoed(t, "aes") +} + +func TestTestRequestEchoesLongPayloadChaChaPoly(t *testing.T) { + assertTestRequestEchoed(t, "chachapoly") +} + +// drainUDPTx empties a control's UDP tx queue without blocking. +func drainUDPTx(c *nebula.Control) { + for c.GetFromUDP(false) != nil { + } +} + +// nextUDPTxOfType returns the next packet a control transmits whose nebula +// header matches (wantType, wantSub), skipping unrelated packets. +// It fails the test if none arrives within the timeout. +func nextUDPTxOfType(t *testing.T, c *nebula.Control, wantType header.MessageType, wantSub header.MessageSubType, within time.Duration) *udp.Packet { + t.Helper() + ch := c.GetUDPTxChan() + timeout := time.After(within) + for { + select { + case p := <-ch: + var h header.H + if err := h.Parse(p.Data); err == nil && h.Type == wantType && h.Subtype == wantSub { + return p + } + case <-timeout: + t.Fatalf("timed out waiting for a %v/%v packet on the udp tx queue", wantType, wantSub) + return nil + } + } +} diff --git a/outside.go b/outside.go index 7ebd4b5e..aad776bf 100644 --- a/outside.go +++ b/outside.go @@ -150,7 +150,8 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, case header.TestReply: // No-op, useful for the Roaming and connectionManager side-effects above case header.TestRequest: - f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, out) + //recycle the input packet ciphertext as our output buffer + f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet) default: hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h) return From 95d98b1f4ba425f40f7943b4f8183f091863e2fc Mon Sep 17 00:00:00 2001 From: Jack Doan Date: Tue, 7 Jul 2026 13:10:11 -0500 Subject: [PATCH 05/25] firewall: move conntrack check after cert+IP verification (#1779) --- firewall.go | 10 ++-- firewall_test.go | 153 +++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 158 insertions(+), 5 deletions(-) diff --git a/firewall.go b/firewall.go index eb120fa6..84c505e7 100644 --- a/firewall.go +++ b/firewall.go @@ -423,11 +423,6 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table") // Drop returns an error if the packet should be dropped, explaining why. It // returns nil if the packet should not be dropped. func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error { - // Check if we spoke to this tuple, if we did then allow this packet - if f.inConns(fp, h, caPool, localCache) { - return nil - } - // Make sure remote address matches nebula certificate, and determine how to treat it if h.networks == nil { // Simple case: Certificate has one address and no unsafe networks @@ -461,6 +456,11 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool * return ErrInvalidLocalIP } + // Check if we spoke to this tuple, if we did then allow this packet + if f.inConns(fp, h, caPool, localCache) { + return nil + } + table := f.OutRules if incoming { table = f.InRules diff --git a/firewall_test.go b/firewall_test.go index 9373f1fd..499f3cc7 100644 --- a/firewall_test.go +++ b/firewall_test.go @@ -916,6 +916,159 @@ func TestFirewall_DropIPSpoofing(t *testing.T) { assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP) } +func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) { + l := test.NewLoggerWithOutput(&bytes.Buffer{}) + + myVpnNetworksTable := new(bart.Lite) + myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24")) + + owner := &dummyCert{ + name: "owner", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")}, + } + + victim := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "victim", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")}, + }, + } + victimHI := HostInfo{ + ConnectionState: &ConnectionState{peerCert: victim}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")}, + } + victimHI.buildNetworks(myVpnNetworksTable, victim.Certificate) + + attacker := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "attacker", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.3/24")}, + }, + } + attackerHI := HostInfo{ + ConnectionState: &ConnectionState{peerCert: attacker}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.3")}, + } + attackerHI.buildNetworks(myVpnNetworksTable, attacker.Certificate) + + fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner) + // Allow any inbound traffic that passes the cert / source-IP checks. + require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + cp := cert.NewCAPool() + + flow := firewall.Packet{ + LocalAddr: netip.MustParseAddr("192.0.2.1"), + RemoteAddr: netip.MustParseAddr("192.0.2.2"), + LocalPort: 443, + RemotePort: 55000, + Protocol: firewall.ProtoUDP, + } + + require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil), + "victim's own traffic from its own overlay IP must be allowed") + + unseen := flow + unseen.RemotePort = 55001 + assert.Equal(t, ErrInvalidRemoteIP, fw.Drop(unseen, true, &attackerHI, cp, nil), + "sanity: attacker forging victim's source IP must be rejected when no conntrack entry exists") + + got := fw.Drop(flow, true, &attackerHI, cp, nil) + t.Logf("attacker replaying victim's 4-tuple: Drop returned %v (nil == packet ALLOWED == spoof succeeded)", got) + assert.Equal(t, ErrInvalidRemoteIP, got, + "SECURITY: attacker spoofed victim's overlay source IP (192.0.2.2) by reusing an existing conntrack 4-tuple; Drop returned %v instead of rejecting", got) +} + +// BenchmarkFirewallDropConntrackHit measures Drop on an already-established flow +// (a conntrack hit). This is the fast path that the source-IP<->cert binding +// reordering adds work to, so it quantifies the cost of moving the address checks +// ahead of the conntrack lookup. Cases: +// - simple: peer cert has one address, no unsafe networks (h.networks == nil), +// so the remote-address check is a single netip.Addr compare. +// - complex: peer cert has unsafe networks (h.networks populated), so the +// remote-address check is a BART lookup. +// - noCache/localCache: whether a per-batch ConntrackCache is supplied, which in +// the original code let the fast path skip straight past the address checks. +func BenchmarkFirewallDropConntrackHit(b *testing.B) { + l := test.NewLoggerWithOutput(&bytes.Buffer{}) + + myVpnNetworksTable := new(bart.Lite) + myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24")) + + owner := &dummyCert{ + name: "owner", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")}, + } + + simpleCert := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "simple", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")}, + }, + } + simpleHost := &HostInfo{ + ConnectionState: &ConnectionState{peerCert: simpleCert}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")}, + } + simpleHost.buildNetworks(myVpnNetworksTable, simpleCert.Certificate) + + complexCert := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "complex", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")}, + unsafeNetworks: []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}, + }, + } + complexHost := &HostInfo{ + ConnectionState: &ConnectionState{peerCert: complexCert}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")}, + } + complexHost.buildNetworks(myVpnNetworksTable, complexCert.Certificate) + + cp := cert.NewCAPool() + + flow := firewall.Packet{ + LocalAddr: netip.MustParseAddr("192.0.2.1"), + RemoteAddr: netip.MustParseAddr("192.0.2.2"), + LocalPort: 443, + RemotePort: 55000, + Protocol: firewall.ProtoUDP, + } + + cases := []struct { + name string + host *HostInfo + useCache bool + }{ + {"simple/noCache", simpleHost, false}, + {"simple/localCache", simpleHost, true}, + {"complex/noCache", complexHost, false}, + {"complex/localCache", complexHost, true}, + } + + for _, tc := range cases { + b.Run(tc.name, func(b *testing.B) { + fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner) + require.NoError(b, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + + // Establish the conntrack entry so every benchmarked Drop is a hit. + require.NoError(b, fw.Drop(flow, true, tc.host, cp, nil)) + + var cache firewall.ConntrackCache + if tc.useCache { + cache = firewall.ConntrackCache{} + } + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := fw.Drop(flow, true, tc.host, cp, cache); err != nil { + b.Fatal(err) + } + } + }) + } +} + func BenchmarkLookup(b *testing.B) { ml := func(m map[string]struct{}, a [][]string) { for n := 0; n < b.N; n++ { From 6aa3363d85a72b1f978327cd2b765aecf1ab0fda Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 13:16:57 -0500 Subject: [PATCH 06/25] Add explicit unmarshaller for signing and key agreement public keys (#1777) --- cert/pem.go | 34 ++++++++++- cert/pem_test.go | 156 ++++++++++++++++++++++++++--------------------- 2 files changed, 120 insertions(+), 70 deletions(-) diff --git a/cert/pem.go b/cert/pem.go index 84221b22..caa19b11 100644 --- a/cert/pem.go +++ b/cert/pem.go @@ -148,6 +148,9 @@ func MarshalSigningPublicKeyToPEM(curve Curve, b []byte) []byte { } } +// UnmarshalPublicKeyFromPEM will try to unmarshal the first pem block in a byte array, returning any non +// consumed data or an error on failure. Only key-agreement (ECDH) public key banners are accepted. +// Use UnmarshalSigningPublicKeyFromPEM for Ed25519/ECDSA banners. func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) { k, r := pem.Decode(b) if k == nil { @@ -156,10 +159,10 @@ func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) { var expectedLen int var curve Curve switch k.Type { - case X25519PublicKeyBanner, Ed25519PublicKeyBanner: + case X25519PublicKeyBanner: expectedLen = 32 curve = Curve_CURVE25519 - case P256PublicKeyBanner, ECDSAP256PublicKeyBanner: + case P256PublicKeyBanner: // Uncompressed expectedLen = 65 curve = Curve_P256 @@ -172,6 +175,33 @@ func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) { return k.Bytes, r, curve, nil } +// UnmarshalSigningPublicKeyFromPEM will try to unmarshal the first pem block in a byte array, returning any non +// consumed data or an error on failure. Only Ed25519/ECDSA public key banners are accepted. +// Use UnmarshalPublicKeyFromPEM for X25519/P256 (ECDH) banners. +func UnmarshalSigningPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) { + k, r := pem.Decode(b) + if k == nil { + return nil, r, 0, fmt.Errorf("input did not contain a valid PEM encoded block") + } + var expectedLen int + var curve Curve + switch k.Type { + case Ed25519PublicKeyBanner: + expectedLen = 32 + curve = Curve_CURVE25519 + case ECDSAP256PublicKeyBanner: + // Uncompressed + expectedLen = 65 + curve = Curve_P256 + default: + return nil, r, 0, fmt.Errorf("bytes did not contain a proper Ed25519/ECDSA public key banner") + } + if len(k.Bytes) != expectedLen { + return nil, r, 0, fmt.Errorf("key was not %d bytes, is invalid %s public key", expectedLen, curve) + } + return k.Bytes, r, curve, nil +} + func MarshalPrivateKeyToPEM(curve Curve, b []byte) []byte { switch curve { case Curve_CURVE25519: diff --git a/cert/pem_test.go b/cert/pem_test.go index ff623541..6012dab3 100644 --- a/cert/pem_test.go +++ b/cert/pem_test.go @@ -255,60 +255,6 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= func TestUnmarshalPublicKeyFromPEM(t *testing.T) { t.Parallel() pubKey := []byte(`# A good key ------BEGIN NEBULA ED25519 PUBLIC KEY----- -AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= ------END NEBULA ED25519 PUBLIC KEY----- -`) - shortKey := []byte(`# A short key ------BEGIN NEBULA ED25519 PUBLIC KEY----- -AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA== ------END NEBULA ED25519 PUBLIC KEY----- -`) - invalidBanner := []byte(`# Invalid banner ------BEGIN NOT A NEBULA PUBLIC KEY----- -AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= ------END NOT A NEBULA PUBLIC KEY----- -`) - invalidPem := []byte(`# Not a valid PEM format --BEGIN NEBULA ED25519 PUBLIC KEY----- -AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= --END NEBULA ED25519 PUBLIC KEY-----`) - - keyBundle := appendByteSlices(pubKey, shortKey, invalidBanner, invalidPem) - - // Success test case - k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle) - assert.Len(t, k, 32) - assert.Equal(t, Curve_CURVE25519, curve) - require.NoError(t, err) - assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem)) - - // Fail due to short key - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) - assert.Nil(t, k) - assert.Equal(t, Curve_CURVE25519, curve) - assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem)) - require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key") - - // Fail due to invalid banner - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) - assert.Nil(t, k) - assert.Equal(t, Curve_CURVE25519, curve) - require.EqualError(t, err, "bytes did not contain a proper public key banner") - assert.Equal(t, rest, invalidPem) - - // Fail due to invalid PEM format, because - // it's missing the requisite pre-encapsulation boundary. - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) - assert.Nil(t, k) - assert.Equal(t, Curve_CURVE25519, curve) - assert.Equal(t, rest, invalidPem) - require.EqualError(t, err, "input did not contain a valid PEM encoded block") -} - -func TestUnmarshalX25519PublicKey(t *testing.T) { - t.Parallel() - pubKey := []byte(`# A good key -----BEGIN NEBULA X25519 PUBLIC KEY----- AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= -----END NEBULA X25519 PUBLIC KEY----- @@ -319,7 +265,7 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA AAAAAAAAAAAAAAAAAAAAAAA= -----END NEBULA P256 PUBLIC KEY----- `) - oldPubP256Key := []byte(`# A good key + signingKey := []byte(`# A signing key has the wrong scope for this function -----BEGIN NEBULA ECDSA P256 PUBLIC KEY----- AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA AAAAAAAAAAAAAAAAAAAAAAA= @@ -340,44 +286,118 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= -END NEBULA X25519 PUBLIC KEY-----`) - keyBundle := appendByteSlices(pubKey, pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem) + keyBundle := appendByteSlices(pubKey, pubP256Key, signingKey, shortKey, invalidBanner, invalidPem) - // Success test case + // X25519 key k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle) assert.Len(t, k, 32) require.NoError(t, err) - assert.Equal(t, rest, appendByteSlices(pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem)) + assert.Equal(t, rest, appendByteSlices(pubP256Key, signingKey, shortKey, invalidBanner, invalidPem)) assert.Equal(t, Curve_CURVE25519, curve) - // Success test case + // P256 key k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) assert.Len(t, k, 65) require.NoError(t, err) - assert.Equal(t, rest, appendByteSlices(oldPubP256Key, shortKey, invalidBanner, invalidPem)) + assert.Equal(t, rest, appendByteSlices(signingKey, shortKey, invalidBanner, invalidPem)) assert.Equal(t, Curve_P256, curve) - // Success test case - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) - assert.Len(t, k, 65) - require.NoError(t, err) + // Reject a signing public key (Ed25519/ECDSA banner) + k, rest, _, err = UnmarshalPublicKeyFromPEM(rest) + assert.Nil(t, k) assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem)) - assert.Equal(t, Curve_P256, curve) + require.EqualError(t, err, "bytes did not contain a proper public key banner") // Fail due to short key - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) + k, rest, _, err = UnmarshalPublicKeyFromPEM(rest) assert.Nil(t, k) assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem)) require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key") // Fail due to invalid banner - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) + k, rest, _, err = UnmarshalPublicKeyFromPEM(rest) assert.Nil(t, k) require.EqualError(t, err, "bytes did not contain a proper public key banner") assert.Equal(t, rest, invalidPem) // Fail due to invalid PEM format, because // it's missing the requisite pre-encapsulation boundary. - k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) + k, rest, _, err = UnmarshalPublicKeyFromPEM(rest) + assert.Nil(t, k) + assert.Equal(t, rest, invalidPem) + require.EqualError(t, err, "input did not contain a valid PEM encoded block") +} + +func TestUnmarshalSigningPublicKeyFromPEM(t *testing.T) { + t.Parallel() + pubKey := []byte(`# A good key +-----BEGIN NEBULA ED25519 PUBLIC KEY----- +AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= +-----END NEBULA ED25519 PUBLIC KEY----- +`) + pubP256Key := []byte(`# A good key +-----BEGIN NEBULA ECDSA P256 PUBLIC KEY----- +AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA +AAAAAAAAAAAAAAAAAAAAAAA= +-----END NEBULA ECDSA P256 PUBLIC KEY----- +`) + ecdhKey := []byte(`# A key-agreement key has the wrong scope for this function +-----BEGIN NEBULA X25519 PUBLIC KEY----- +AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= +-----END NEBULA X25519 PUBLIC KEY----- +`) + shortKey := []byte(`# A short key +-----BEGIN NEBULA ED25519 PUBLIC KEY----- +AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA== +-----END NEBULA ED25519 PUBLIC KEY----- +`) + invalidBanner := []byte(`# Invalid banner +-----BEGIN NOT A NEBULA PUBLIC KEY----- +AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= +-----END NOT A NEBULA PUBLIC KEY----- +`) + invalidPem := []byte(`# Not a valid PEM format +-BEGIN NEBULA ED25519 PUBLIC KEY----- +AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= +-END NEBULA ED25519 PUBLIC KEY-----`) + + keyBundle := appendByteSlices(pubKey, pubP256Key, ecdhKey, shortKey, invalidBanner, invalidPem) + + // Ed25519 key + k, rest, curve, err := UnmarshalSigningPublicKeyFromPEM(keyBundle) + assert.Len(t, k, 32) + require.NoError(t, err) + assert.Equal(t, rest, appendByteSlices(pubP256Key, ecdhKey, shortKey, invalidBanner, invalidPem)) + assert.Equal(t, Curve_CURVE25519, curve) + + // ECDSA P256 key + k, rest, curve, err = UnmarshalSigningPublicKeyFromPEM(rest) + assert.Len(t, k, 65) + require.NoError(t, err) + assert.Equal(t, rest, appendByteSlices(ecdhKey, shortKey, invalidBanner, invalidPem)) + assert.Equal(t, Curve_P256, curve) + + // Reject a key-agreement public key (X25519/P256 banner) + k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) + assert.Nil(t, k) + assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem)) + require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA public key banner") + + // Fail due to short key + k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) + assert.Nil(t, k) + assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem)) + require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key") + + // Fail due to invalid banner + k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) + assert.Nil(t, k) + require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA public key banner") + assert.Equal(t, rest, invalidPem) + + // Fail due to invalid PEM format, because + // it's missing the requisite pre-encapsulation boundary. + k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) assert.Nil(t, k) assert.Equal(t, rest, invalidPem) require.EqualError(t, err, "input did not contain a valid PEM encoded block") From 0a953915bb7277ced88d926255935a9edc2f42a1 Mon Sep 17 00:00:00 2001 From: John Maguire Date: Tue, 7 Jul 2026 18:27:53 +0000 Subject: [PATCH 07/25] Make HostInfo.remote atomic to fix torn reads on the send path (#1773) --- control.go | 4 ++-- control_test.go | 14 ++++++++------ hostmap.go | 17 ++++++++++++----- inside.go | 8 ++++---- outside.go | 14 ++++++++------ punchy.go | 4 ++-- relay_manager.go | 6 +++--- 7 files changed, 39 insertions(+), 28 deletions(-) diff --git a/control.go b/control.go index ef58988b..053feab5 100644 --- a/control.go +++ b/control.go @@ -305,7 +305,7 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) { c.l.Debug("Sending close tunnel message", "vpnAddrs", h.vpnAddrs, - "udpAddr", h.remote, + "udpAddr", h.GetRemote(), ) closed++ } @@ -350,7 +350,7 @@ func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo { RemoteAddrs: h.remotes.CopyAddrs(preferredRanges), CurrentRelaysToMe: h.relayState.CopyRelayIps(), CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(), - CurrentRemote: h.remote, + CurrentRemote: h.GetRemote(), } for i, a := range h.vpnAddrs { diff --git a/control_test.go b/control_test.go index 5e381c46..dae759a3 100644 --- a/control_test.go +++ b/control_test.go @@ -42,8 +42,7 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) { assert.True(t, ok) crt := &dummyCert{} - hm.unlockedAddHostInfo(&HostInfo{ - remote: remote1, + hi := &HostInfo{ remotes: remotes, ConnectionState: &ConnectionState{ peerCert: &cert.CachedCertificate{Certificate: crt}, @@ -56,13 +55,14 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) { relayForByAddr: map[netip.Addr]*Relay{}, relayForByIdx: map[uint32]*Relay{}, }, - }, &Interface{}) + } + hi.remote.Store(&remote1) + hm.unlockedAddHostInfo(hi, &Interface{}) vpnIp2, ok := netip.AddrFromSlice(ipNet2.IP) assert.True(t, ok) - hm.unlockedAddHostInfo(&HostInfo{ - remote: remote1, + hi2 := &HostInfo{ remotes: remotes, ConnectionState: &ConnectionState{ peerCert: nil, @@ -75,7 +75,9 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) { relayForByAddr: map[netip.Addr]*Relay{}, relayForByIdx: map[uint32]*Relay{}, }, - }, &Interface{}) + } + hi2.remote.Store(&remote1) + hm.unlockedAddHostInfo(hi2, &Interface{}) c := Control{ state: StateReady, diff --git a/hostmap.go b/hostmap.go index 957894b6..e7dd17a0 100644 --- a/hostmap.go +++ b/hostmap.go @@ -229,7 +229,7 @@ const ( ) type HostInfo struct { - remote netip.AddrPort + remote atomic.Pointer[netip.AddrPort] remotes *RemoteList promoteCounter atomic.Uint32 ConnectionState *ConnectionState @@ -684,7 +684,7 @@ func (hm *HostMap) ForEachIndex(f controlEach) { func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interface) { c := i.promoteCounter.Add(1) if c%ifce.tryPromoteEvery.Load() == 0 { - remote := i.remote + remote := i.GetRemote() // return early if we are already on a preferred remote if remote.IsValid() { @@ -726,11 +726,18 @@ func (i *HostInfo) GetCert() *cert.CachedCertificate { return nil } +func (i *HostInfo) GetRemote() netip.AddrPort { + if p := i.remote.Load(); p != nil { + return *p + } + return netip.AddrPort{} +} + // TODO: Maybe use ViaSender here? func (i *HostInfo) SetRemote(remote netip.AddrPort) { // We copy here because we likely got this remote from a source that reuses the object - if i.remote != remote { - i.remote = remote + if i.GetRemote() != remote { + i.remote.Store(&remote) i.remotes.LearnRemote(i.vpnAddrs[0], remote) } } @@ -742,7 +749,7 @@ func (i *HostInfo) SetRemoteIfPreferred(hm *HostMap, via ViaSender) bool { return false } - currentRemote := i.remote + currentRemote := i.GetRemote() if !currentRemote.IsValid() { i.SetRemote(via.UdpAddr) return true diff --git a/inside.go b/inside.go index 27a6f758..fc079ba4 100644 --- a/inside.go +++ b/inside.go @@ -333,7 +333,7 @@ func (f *Interface) SendVia(via *HostInfo, via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err) return } - err = f.writers[0].WriteTo(out, via.remote) + err = f.writers[0].WriteTo(out, via.GetRemote()) if err != nil { via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err) } @@ -344,7 +344,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType if ci.eKey == nil { return } - useRelay := !remote.IsValid() && !hostinfo.remote.IsValid() + useRelay := !remote.IsValid() && !hostinfo.GetRemote().IsValid() fullOut := out if useRelay { @@ -403,8 +403,8 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType "udpAddr", remote, ) } - } else if hostinfo.remote.IsValid() { - err = f.writers[q].WriteTo(out, hostinfo.remote) + } else if hr := hostinfo.GetRemote(); hr.IsValid() { + err = f.writers[q].WriteTo(out, hr) if err != nil { hostinfo.logger(f.l).Error("Failed to write outgoing packet", "error", err, diff --git a/outside.go b/outside.go index aad776bf..29607eae 100644 --- a/outside.go +++ b/outside.go @@ -277,7 +277,8 @@ func (f *Interface) sendCloseTunnel(h *HostInfo) { } func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) { - if !via.IsRelayed && hostinfo.remote != via.UdpAddr { + curRemote := hostinfo.GetRemote() + if !via.IsRelayed && curRemote != via.UdpAddr { if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) { if f.l.Enabled(context.Background(), slog.LevelDebug) { hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr) @@ -289,7 +290,7 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) { if f.l.Enabled(context.Background(), slog.LevelDebug) { hostinfo.logger(f.l).Debug("Suppressing roam back to previous remote", "suppressSeconds", RoamingSuppressSeconds, - "udpAddr", hostinfo.remote, + "udpAddr", curRemote, "newAddr", via.UdpAddr, ) } @@ -297,11 +298,11 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) { } hostinfo.logger(f.l).Info("Host roamed to new udp ip/port.", - "udpAddr", hostinfo.remote, + "udpAddr", curRemote, "newAddr", via.UdpAddr, ) hostinfo.lastRoam = time.Now() - hostinfo.lastRoamRemote = hostinfo.remote + hostinfo.lastRoamRemote = curRemote hostinfo.SetRemote(via.UdpAddr) } @@ -590,10 +591,11 @@ func (f *Interface) handleRecvError(addr netip.AddrPort, h *header.H) { return } - if hostinfo.remote.IsValid() && hostinfo.remote != addr { + hr := hostinfo.GetRemote() + if hr.IsValid() && hr != addr { f.l.Info("Someone spoofing recv_errors?", "addr", addr, - "hostinfoRemote", hostinfo.remote, + "hostinfoRemote", hr, ) return } diff --git a/punchy.go b/punchy.go index 38a0e1ca..4bce4392 100644 --- a/punchy.go +++ b/punchy.go @@ -174,9 +174,9 @@ func (p *Punchy) SendPunch(hostinfo *HostInfo) { if p.punchEverything.Load() { p.sendPunchToAllRemotes(hostinfo) - } else if hostinfo.remote.IsValid() { + } else if hr := hostinfo.GetRemote(); hr.IsValid() { p.metricPunchyTx.Inc(1) - p.punchConn.WriteTo([]byte{1}, hostinfo.remote) + p.punchConn.WriteTo([]byte{1}, hr) } } diff --git a/relay_manager.go b/relay_manager.go index 985225f4..3d396883 100644 --- a/relay_manager.go +++ b/relay_manager.go @@ -94,7 +94,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak } relayHostInfo := rm.hostmap.QueryVpnAddr(relay) - if relayHostInfo == nil || !relayHostInfo.remote.IsValid() { + if relayHostInfo == nil || !relayHostInfo.GetRemote().IsValid() { hl.Log(context.Background(), level, "Establish tunnel to relay target", "relay", relay.String()) f.Handshake(relay) continue @@ -104,7 +104,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp) if !ok { // No relays exist or requested yet. - if relayHostInfo.remote.IsValid() { + if relayHostInfo.GetRemote().IsValid() { idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested) if err != nil { hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err) @@ -508,7 +508,7 @@ func (rm *relayManager) handleCreateRelayRequest(v cert.Version, h *HostInfo, f f.Handshake(target) return } - if !peer.remote.IsValid() { + if !peer.GetRemote().IsValid() { // Only create relays to peers for whom I have a direct connection return } From abfeb502a8fcb45e83ab8d09d4a917974a665c4b Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 13:28:52 -0500 Subject: [PATCH 08/25] iputil: fix infinite loop in ipv6FindUpperProtocol from uint8 extension-header length overflow (#1784) --- iputil/packet.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/iputil/packet.go b/iputil/packet.go index 99893822..c0c1921e 100644 --- a/iputil/packet.go +++ b/iputil/packet.go @@ -344,7 +344,7 @@ func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm return nextHeader, offset, isFragment } nextHeader = packet[offset] - offset += int(packet[offset+1]+1) << 3 + offset += (int(packet[offset+1]) + 1) << 3 case 44: // Fragment if len(packet) < offset+8 { @@ -361,7 +361,7 @@ func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragm return nextHeader, offset, isFragment } nextHeader = packet[offset] - offset += int(packet[offset+1]+2) << 2 + offset += (int(packet[offset+1]) + 2) << 2 default: return nextHeader, offset, isFragment From 647775d8c35b184fe2f96a425e78d02f40ed5e41 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Tue, 7 Jul 2026 13:40:12 -0500 Subject: [PATCH 09/25] Bump golang.zx2c4.com/wireguard/windows (#1665) Bumps the zx2c4-dependencies group with 1 update in the / directory: golang.zx2c4.com/wireguard/windows. Updates `golang.zx2c4.com/wireguard/windows` from 0.6.1 to 1.0.1 --- updated-dependencies: - dependency-name: golang.zx2c4.com/wireguard/windows dependency-version: 1.0.1 dependency-type: direct:production update-type: version-update:semver-major dependency-group: zx2c4-dependencies ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- go.mod | 6 +++--- go.sum | 12 ++++++------ 2 files changed, 9 insertions(+), 9 deletions(-) diff --git a/go.mod b/go.mod index 40260040..7a66fb10 100644 --- a/go.mod +++ b/go.mod @@ -32,7 +32,7 @@ require ( golang.org/x/term v0.44.0 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b - golang.zx2c4.com/wireguard/windows v0.6.1 + golang.zx2c4.com/wireguard/windows v1.0.1 google.golang.org/protobuf v1.36.11 gopkg.in/yaml.v3 v3.0.1 gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe @@ -50,7 +50,7 @@ require ( github.com/prometheus/procfs v0.16.1 // indirect github.com/vishvananda/netns v0.0.5 // indirect go.yaml.in/yaml/v2 v2.4.2 // indirect - golang.org/x/mod v0.34.0 // indirect + golang.org/x/mod v0.36.0 // indirect golang.org/x/time v0.5.0 // indirect - golang.org/x/tools v0.43.0 // indirect + golang.org/x/tools v0.45.0 // indirect ) diff --git a/go.sum b/go.sum index 0555a9db..e178d162 100644 --- a/go.sum +++ b/go.sum @@ -170,8 +170,8 @@ golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPI golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg= golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= -golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI= -golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY= +golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= +golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= @@ -223,8 +223,8 @@ golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtn golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28= golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE= golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA= -golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s= -golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0= +golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= +golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= @@ -233,8 +233,8 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI= golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo= golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4= -golang.zx2c4.com/wireguard/windows v0.6.1 h1:XMaKojH1Hs/raMrmnir4n35nTvzvWj7NmSYzHn2F4qU= -golang.zx2c4.com/wireguard/windows v0.6.1/go.mod h1:04aqInu5GYuTFvMuDw/rKBAF7mHrltW/3rekpfbbZDM= +golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8= +golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs= google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= From 19ad3bb9046acf6659d1c63b651f4913695046f9 Mon Sep 17 00:00:00 2001 From: Jack Doan Date: Tue, 7 Jul 2026 14:41:01 -0500 Subject: [PATCH 10/25] correctly discard nil proto addresses (#1785) --- control_test.go | 153 +++++++++++++++++++++++++++++++++++++++++++++++ lighthouse.go | 10 +++- relay_manager.go | 18 ++++++ remote_list.go | 12 ++++ 4 files changed, 192 insertions(+), 1 deletion(-) diff --git a/control_test.go b/control_test.go index dae759a3..94ee4ee3 100644 --- a/control_test.go +++ b/control_test.go @@ -1,6 +1,8 @@ package nebula import ( + "bytes" + "log/slog" "net" "net/netip" "reflect" @@ -9,6 +11,7 @@ import ( "github.com/slackhq/nebula/cert" "github.com/slackhq/nebula/test" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) func TestControl_GetHostInfoByVpnIp(t *testing.T) { @@ -121,3 +124,153 @@ func assertFields(t *testing.T, expected []string, actualStruct any) { assert.Equal(t, expected, fields) } + +// alwaysAllowV4/V6 are check funcs that accept every entry (including nil pointers), +// letting us inject a nil *V4AddrPort/*V6AddrPort into a RemoteList's reported cache +// the same way a malformed proto message off the wire could. +func alwaysAllowV4(netip.Addr, *V4AddrPort) bool { return true } +func alwaysAllowV6(netip.Addr, *V6AddrPort) bool { return true } + +// TestGetRelays_SkipsNilRelayAddrs proves GetRelays tolerates nil entries in the +// RelayVpnAddrs proto slice (which protoAddrToNetAddr would nil-deref on) and still +// returns the valid relays, including the legacy OldRelayVpnAddrs. +func TestGetRelays_SkipsNilRelayAddrs(t *testing.T) { + good := netip.MustParseAddr("10.0.0.9") + + d := &NebulaMetaDetails{ + OldRelayVpnAddrs: []uint32{0x0a000001}, // 10.0.0.1 + RelayVpnAddrs: []*Addr{ + nil, + netAddrToProtoAddr(good), + nil, + }, + } + + var relays []netip.Addr + require.NotPanics(t, func() { relays = d.GetRelays() }) + + assert.Equal(t, []netip.Addr{ + netip.MustParseAddr("10.0.0.1"), + good, + }, relays) +} + +// TestGetRelays_AllNil ensures an all-nil RelayVpnAddrs slice yields no relays and no panic. +func TestGetRelays_AllNil(t *testing.T) { + d := &NebulaMetaDetails{RelayVpnAddrs: []*Addr{nil, nil}} + var relays []netip.Addr + require.NotPanics(t, func() { relays = d.GetRelays() }) + assert.Empty(t, relays) +} + +// TestRemoteList_CopyCache_SkipsNilReported proves CopyCache skips nil reported +// pointers (v4 and v6) instead of nil-dereferencing them in protoV*AddrPortToNetAddrPort. +func TestRemoteList_CopyCache_SkipsNilReported(t *testing.T) { + owner := netip.MustParseAddr("10.0.0.1") + rl := NewRemoteList([]netip.Addr{owner}, nil) + + rl.unlockedSetV4(owner, owner, []*V4AddrPort{ + nil, + newIp4AndPortFromString("1.2.3.4:5"), + nil, + }, alwaysAllowV4) + + rl.unlockedSetV6(owner, owner, []*V6AddrPort{ + nil, + newIp6AndPortFromString("[1::1]:6"), + nil, + }, alwaysAllowV6) + + var cm *CacheMap + require.NotPanics(t, func() { cm = rl.CopyCache() }) + + c := (*cm)[owner.String()] + require.NotNil(t, c) + assert.ElementsMatch(t, []netip.AddrPort{ + netip.MustParseAddrPort("1.2.3.4:5"), + netip.MustParseAddrPort("[1::1]:6"), + }, c.Reported) +} + +// TestRemoteList_Rebuild_SkipsNilReported drives unlockedCollect (via Rebuild) with +// nil reported entries and confirms only the valid addresses survive, with no panic. +func TestRemoteList_Rebuild_SkipsNilReported(t *testing.T) { + owner := netip.MustParseAddr("10.0.0.1") + rl := NewRemoteList([]netip.Addr{owner}, nil) + + rl.unlockedSetV4(owner, owner, []*V4AddrPort{ + nil, + newIp4AndPortFromString("1.2.3.4:5"), + }, alwaysAllowV4) + rl.unlockedSetV6(owner, owner, []*V6AddrPort{ + newIp6AndPortFromString("[1::1]:6"), + nil, + }, alwaysAllowV6) + + require.NotPanics(t, func() { rl.Rebuild([]netip.Prefix{}) }) + + assert.ElementsMatch(t, []netip.AddrPort{ + netip.MustParseAddrPort("1.2.3.4:5"), + netip.MustParseAddrPort("[1::1]:6"), + }, rl.addrs) +} + +// newRelayControl marshals a NebulaControl the way it arrives on the wire so we can feed +// it through HandleControlMsg's unmarshal + validate path. +func newRelayControl(t *testing.T, typ NebulaControl_MessageType, from, to *Addr) []byte { + t.Helper() + msg := &NebulaControl{ + Type: typ, + RelayFromAddr: from, + RelayToAddr: to, + } + b, err := msg.Marshal() + require.NoError(t, err) + return b +} + +// TestRelayManager_HandleControlMsg_NilRelayAddrs verifies the validation block added to +// HandleControlMsg: CreateRelay{Request,Response} carrying a nil RelayFromAddr or +// RelayToAddr are dropped with a debug log rather than nil-dereferencing downstream. +func TestRelayManager_HandleControlMsg_NilRelayAddrs(t *testing.T) { + good := netAddrToProtoAddr(netip.MustParseAddr("10.0.0.9")) + + cases := []struct { + name string + typ NebulaControl_MessageType + from *Addr + to *Addr + wantLog string // debug substring expected, "" == expect no drop log + }{ + {"request nil from", NebulaControl_CreateRelayRequest, nil, good, "nil RelayFromAddr"}, + {"request nil to", NebulaControl_CreateRelayRequest, good, nil, "nil RelayToAddr"}, + {"request both nil", NebulaControl_CreateRelayRequest, nil, nil, "nil RelayFromAddr"}, + {"response nil from", NebulaControl_CreateRelayResponse, nil, good, "nil RelayFromAddr"}, + {"response nil to", NebulaControl_CreateRelayResponse, good, nil, "nil RelayToAddr"}, + // A non-relay control type is not subject to the relay-addr validation and must + // pass through it untouched (the final switch simply no-ops on it). + {"unrelated type nil addrs", NebulaControl_None, nil, nil, ""}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + var buf bytes.Buffer + l := test.NewLoggerWithOutputAndLevel(&buf, slog.LevelDebug) + rm := &relayManager{l: l, hostmap: newHostMap(l)} + rm.useRelays.Store(true) + + f := &Interface{l: l} + h := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("10.0.0.2")}, localIndexId: 1} + + d := newRelayControl(t, tc.typ, tc.from, tc.to) + + require.NotPanics(t, func() { rm.HandleControlMsg(h, d, f) }) + + if tc.wantLog == "" { + assert.NotContains(t, buf.String(), "nil Relay") + } else { + assert.Contains(t, buf.String(), tc.wantLog) + } + }) + } +} diff --git a/lighthouse.go b/lighthouse.go index d23e84b8..5557b1c4 100644 --- a/lighthouse.go +++ b/lighthouse.go @@ -1418,6 +1418,9 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn remoteAllowList := lhh.lh.GetRemoteAllowList() for _, a := range n.Details.V4AddrPorts { + if a == nil { + continue + } b := protoV4AddrPortToNetAddrPort(a) if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) { lhh.lh.punchy.Schedule(b, detailsVpnAddr) @@ -1425,6 +1428,9 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn } for _, a := range n.Details.V6AddrPorts { + if a == nil { + continue + } b := protoV6AddrPortToNetAddrPort(a) if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) { lhh.lh.punchy.Schedule(b, detailsVpnAddr) @@ -1494,7 +1500,9 @@ func (d *NebulaMetaDetails) GetRelays() []netip.Addr { if len(d.RelayVpnAddrs) > 0 { for _, r := range d.RelayVpnAddrs { - relays = append(relays, protoAddrToNetAddr(r)) + if r != nil { + relays = append(relays, protoAddrToNetAddr(r)) + } } } return relays diff --git a/relay_manager.go b/relay_manager.go index 3d396883..318a9f1a 100644 --- a/relay_manager.go +++ b/relay_manager.go @@ -309,6 +309,22 @@ func (rm *relayManager) HandleControlMsg(h *HostInfo, d []byte, f *Interface) { v = cert.Version2 } + // validate: + switch msg.Type { + case NebulaControl_CreateRelayRequest, NebulaControl_CreateRelayResponse: + if msg.RelayFromAddr == nil { + if f.l.Enabled(context.Background(), slog.LevelDebug) { + h.logger(f.l).Debug("Control message received with nil RelayFromAddr", "type", msg.Type) + } + return + } else if msg.RelayToAddr == nil { + if f.l.Enabled(context.Background(), slog.LevelDebug) { + h.logger(f.l).Debug("Control message received with nil RelayToAddr", "type", msg.Type) + } + return + } + } + switch msg.Type { case NebulaControl_CreateRelayRequest: rm.handleCreateRelayRequest(v, h, f, msg) @@ -318,6 +334,7 @@ func (rm *relayManager) HandleControlMsg(h *HostInfo, d []byte, f *Interface) { } func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f *Interface, m *NebulaControl) { + //nil-checks for protoAddrToNetAddr handled by caller relayFrom := protoAddrToNetAddr(m.RelayFromAddr) relayTo := protoAddrToNetAddr(m.RelayToAddr) rm.l.Info("handleCreateRelayResponse", @@ -399,6 +416,7 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f } func (rm *relayManager) handleCreateRelayRequest(v cert.Version, h *HostInfo, f *Interface, m *NebulaControl) { + //nil-checks for protoAddrToNetAddr handled by caller from := protoAddrToNetAddr(m.RelayFromAddr) target := protoAddrToNetAddr(m.RelayToAddr) diff --git a/remote_list.go b/remote_list.go index ef6eb794..9d1b387e 100644 --- a/remote_list.go +++ b/remote_list.go @@ -344,6 +344,9 @@ func (r *RemoteList) CopyCache() *CacheMap { } for _, a := range mc.v4.reported { + if a == nil { + continue + } c.Reported = append(c.Reported, protoV4AddrPortToNetAddrPort(a)) } } @@ -354,6 +357,9 @@ func (r *RemoteList) CopyCache() *CacheMap { } for _, a := range mc.v6.reported { + if a == nil { + continue + } c.Reported = append(c.Reported, protoV6AddrPortToNetAddrPort(a)) } } @@ -582,6 +588,9 @@ func (r *RemoteList) unlockedCollect() { } for _, v := range c.v4.reported { + if v == nil { + continue + } u := protoV4AddrPortToNetAddrPort(v) if !r.unlockedIsBad(u) { addrs = append(addrs, u) @@ -598,6 +607,9 @@ func (r *RemoteList) unlockedCollect() { } for _, v := range c.v6.reported { + if v == nil { + continue + } u := protoV6AddrPortToNetAddrPort(v) if !r.unlockedIsBad(u) { addrs = append(addrs, u) From 32149f3a9388d9eaacb05ed29456c6b8ca211017 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Tue, 7 Jul 2026 14:42:33 -0500 Subject: [PATCH 11/25] Bump github.com/kardianos/service from 1.2.4 to 1.3.0 (#1782) Bumps [github.com/kardianos/service](https://github.com/kardianos/service) from 1.2.4 to 1.3.0. - [Commits](https://github.com/kardianos/service/compare/v1.2.4...v1.3.0) --- updated-dependencies: - dependency-name: github.com/kardianos/service dependency-version: 1.3.0 dependency-type: direct:production update-type: version-update:semver-minor ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index 7a66fb10..b30a97a1 100644 --- a/go.mod +++ b/go.mod @@ -12,7 +12,7 @@ require ( github.com/gaissmai/bart v0.28.0 github.com/gogo/protobuf v1.3.2 github.com/google/gopacket v1.1.19 - github.com/kardianos/service v1.2.4 + github.com/kardianos/service v1.3.0 github.com/miekg/dns v1.1.72 github.com/miekg/pkcs11 v1.1.2 github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f diff --git a/go.sum b/go.sum index e178d162..11e72276 100644 --- a/go.sum +++ b/go.sum @@ -66,8 +66,8 @@ github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/ github.com/json-iterator/go v1.1.11/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= github.com/julienschmidt/httprouter v1.2.0/go.mod h1:SYymIcj16QtmaHHD7aYtjjsJG7VTCxuUUipMqKk8s4w= github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM= -github.com/kardianos/service v1.2.4 h1:XNlGtZOYNx2u91urOdg/Kfmc+gfmuIo1Dd3rEi2OgBk= -github.com/kardianos/service v1.2.4/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc= +github.com/kardianos/service v1.3.0 h1:/LGy+xPP2TM+GLTiCZ2di7cy0Jd/qrawlTUfqKYFdTI= +github.com/kardianos/service v1.3.0/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= From 942ee522e08ea19608097884c01d20b6da1049da Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 15:23:39 -0500 Subject: [PATCH 12/25] lighthouse: unmap 4-in-6 addresses in protoV6AddrPortToNetAddrPort so remote_allow_list v4 rules apply (#1786) --- lighthouse.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lighthouse.go b/lighthouse.go index 5557b1c4..3df74c39 100644 --- a/lighthouse.go +++ b/lighthouse.go @@ -1460,7 +1460,7 @@ func protoV6AddrPortToNetAddrPort(ap *V6AddrPort) netip.AddrPort { b := [16]byte{} binary.BigEndian.PutUint64(b[:8], ap.Hi) binary.BigEndian.PutUint64(b[8:], ap.Lo) - return netip.AddrPortFrom(netip.AddrFrom16(b), uint16(ap.Port)) + return netip.AddrPortFrom(netip.AddrFrom16(b).Unmap(), uint16(ap.Port)) } func netAddrToProtoAddr(addr netip.Addr) *Addr { From 7bd0bc285a9c192416fd554120170a3853182e71 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 15:50:06 -0500 Subject: [PATCH 13/25] sshd: guard trustedKeys/trustedCAs with a mutex to fix a concurrent map crash on reload (#1787) --- sshd/server.go | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/sshd/server.go b/sshd/server.go index 86c52961..e0ee9364 100644 --- a/sshd/server.go +++ b/sshd/server.go @@ -7,6 +7,7 @@ import ( "fmt" "log/slog" "net" + "sync" "github.com/armon/go-radix" "golang.org/x/crypto/ssh" @@ -18,6 +19,8 @@ type SSHServer struct { certChecker *ssh.CertChecker + // authLock guards trustedKeys and trustedCAs + authLock sync.RWMutex // Map of user -> authorized keys trustedKeys map[string]map[string]bool trustedCAs []ssh.PublicKey @@ -45,6 +48,8 @@ func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) { cc := ssh.CertChecker{ IsUserAuthority: func(auth ssh.PublicKey) bool { + s.authLock.RLock() + defer s.authLock.RUnlock() for _, ca := range s.trustedCAs { if bytes.Equal(ca.Marshal(), auth.Marshal()) { return true @@ -57,6 +62,8 @@ func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) { pk := string(pubKey.Marshal()) fp := ssh.FingerprintSHA256(pubKey) + s.authLock.RLock() + defer s.authLock.RUnlock() tk, ok := s.trustedKeys[c.User()] if !ok { return nil, fmt.Errorf("unknown user %s", c.User()) @@ -105,11 +112,15 @@ func (s *SSHServer) SetHostKey(hostPrivateKey []byte) error { } func (s *SSHServer) ClearTrustedCAs() { + s.authLock.Lock() s.trustedCAs = []ssh.PublicKey{} + s.authLock.Unlock() } func (s *SSHServer) ClearAuthorizedKeys() { + s.authLock.Lock() s.trustedKeys = make(map[string]map[string]bool) + s.authLock.Unlock() } // AddTrustedCA adds a trusted CA for user certificates @@ -119,7 +130,9 @@ func (s *SSHServer) AddTrustedCA(pubKey string) error { return err } + s.authLock.Lock() s.trustedCAs = append(s.trustedCAs, pk) + s.authLock.Unlock() s.l.Info("Trusted CA key", "sshKey", pubKey) return nil } @@ -131,6 +144,7 @@ func (s *SSHServer) AddAuthorizedKey(user, pubKey string) error { return err } + s.authLock.Lock() tk, ok := s.trustedKeys[user] if !ok { tk = make(map[string]bool) @@ -138,6 +152,7 @@ func (s *SSHServer) AddAuthorizedKey(user, pubKey string) error { } tk[string(pk.Marshal())] = true + s.authLock.Unlock() s.l.Info("Authorized ssh key", "sshKey", pubKey, "sshUser", user, From e5c0fdad8daab7f1089817f1a0423afb05047645 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 17:05:12 -0500 Subject: [PATCH 14/25] Darwin and openbsd in line with the other bsds for tun support (#1703) --- overlay/tun_darwin.go | 122 ++++++++++++++++++++++++++++++----------- overlay/tun_openbsd.go | 105 +++++++++++++++++++++++++++-------- 2 files changed, 172 insertions(+), 55 deletions(-) diff --git a/overlay/tun_darwin.go b/overlay/tun_darwin.go index 524ef0cd..d30148b9 100644 --- a/overlay/tun_darwin.go +++ b/overlay/tun_darwin.go @@ -23,7 +23,7 @@ import ( ) type tun struct { - io.ReadWriteCloser + f *os.File Device string vpnNetworks []netip.Prefix DefaultMTU int @@ -31,9 +31,6 @@ type tun struct { routeTree atomic.Pointer[bart.Table[routing.Gateways]] linkAddr *netroute.LinkAddr l *slog.Logger - - // cache out buffer since we need to prepend 4 bytes for tun metadata - out []byte } type ifReq struct { @@ -124,11 +121,11 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t } t := &tun{ - ReadWriteCloser: os.NewFile(uintptr(fd), ""), - Device: name, - vpnNetworks: vpnNetworks, - DefaultMTU: c.GetInt("tun.mtu", DefaultMTU), - l: l, + f: os.NewFile(uintptr(fd), ""), + Device: name, + vpnNetworks: vpnNetworks, + DefaultMTU: c.GetInt("tun.mtu", DefaultMTU), + l: l, } err = t.reload(c, true) @@ -158,8 +155,8 @@ func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, e } func (t *tun) Close() error { - if t.ReadWriteCloser != nil { - return t.ReadWriteCloser.Close() + if t.f != nil { + return t.f.Close() } return nil } @@ -502,42 +499,103 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error { return nil } +// tunWritev and tunReadv are linkname'd to x/sys/unix's libc-routed writev/readv stubs so the +// calls go through libSystem's pinned trampoline. A raw syscall.Syscall(SYS_WRITEV/SYS_READV, ...) +// on darwin/arm64 emits an SVC #0x80 trap (see $GOROOT/src/syscall/asm_darwin_arm64.s), the path +// Apple keeps warning they will eventually disallow. We pull the low-level stubs instead of calling +// unix.Writev/unix.Readv because those take [][]byte and rebuild the []Iovec every call, which +// heap-allocates the header; linkname'ing the stubs lets us hand them our own stack-allocated +// iovecs. See golang/go#78049. + +//go:linkname tunWritev golang.org/x/sys/unix.writev +//go:noescape +func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error) + +//go:linkname tunReadv golang.org/x/sys/unix.readv +//go:noescape +func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error) + +// Read pulls one IP packet off the utun device, scattering the 4 byte protocol header away from +// the packet so the payload lands directly in to. func (t *tun) Read(to []byte) (int, error) { - buf := make([]byte, len(to)+4) + var head [4]byte - n, err := t.ReadWriteCloser.Read(buf) + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - copy(to, buf[4:]) - return n - 4, err + var n int + var callErr error + err = rc.Read(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &to[0], Len: uint64(len(to))}, + } + n, callErr = tunReadv(int(fd), iovecs) + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + if n < 4 { + return 0, nil + } + return n - 4, nil } -// Write is only valid for single threaded use +// Write pushes one IP packet onto the utun device. func (t *tun) Write(from []byte) (int, error) { - buf := t.out - if cap(buf) < len(from)+4 { - buf = make([]byte, len(from)+4) - t.out = buf - } - buf = buf[:len(from)+4] - if len(from) == 0 { return 0, syscall.EIO } - // Determine the IP Family for the NULL L2 Header ipVer := from[0] >> 4 - if ipVer == 4 { - buf[3] = syscall.AF_INET - } else if ipVer == 6 { - buf[3] = syscall.AF_INET6 - } else { + var head [4]byte + switch ipVer { + case 4: + head[3] = syscall.AF_INET + case 6: + head[3] = syscall.AF_INET6 + default: return 0, fmt.Errorf("unable to determine IP version from packet") } - copy(buf[4:], from) + // Grab rc as a local so the compiler can devirtualize the call and keep the closure on the stack. + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - n, err := t.ReadWriteCloser.Write(buf) - return n - 4, err + var n int + var callErr error + err = rc.Write(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &from[0], Len: uint64(len(from))}, + } + n, callErr = tunWritev(int(fd), iovecs) + // Type-assert to syscall.Errno so the EAGAIN/EWOULDBLOCK/EINTR check doesn't box the errno + // constants into error interfaces on every call. + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + + return n - 4, nil } func (t *tun) Networks() []netip.Prefix { diff --git a/overlay/tun_openbsd.go b/overlay/tun_openbsd.go index 81362184..41224777 100644 --- a/overlay/tun_openbsd.go +++ b/overlay/tun_openbsd.go @@ -57,8 +57,6 @@ type tun struct { l *slog.Logger f *os.File fd int - // cache out buffer since we need to prepend 4 bytes for tun metadata - out []byte } var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`) @@ -124,42 +122,103 @@ func (t *tun) Close() error { return nil } +// tunWritev and tunReadv are linkname'd to x/sys/unix's libc-routed writev/readv stubs so the +// calls go through libc's pinned trampoline. OpenBSD's pinsyscall protection rejects a raw +// syscall.Syscall(SYS_WRITEV/SYS_READV, ...) because it doesn't originate from a libc-pinned +// address, so we can't use the syscall.Syscall pattern that freebsd / netbsd use. We pull the +// low-level stubs instead of calling unix.Writev/unix.Readv because those take [][]byte and rebuild +// the []Iovec every call, which heap-allocates the header; linkname'ing the stubs lets us hand them +// our own stack-allocated iovecs. See golang/go#78049. + +//go:linkname tunWritev golang.org/x/sys/unix.writev +//go:noescape +func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error) + +//go:linkname tunReadv golang.org/x/sys/unix.readv +//go:noescape +func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error) + +// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from the +// packet so the payload lands directly in to. func (t *tun) Read(to []byte) (int, error) { - buf := make([]byte, len(to)+4) + var head [4]byte - n, err := t.f.Read(buf) + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - copy(to, buf[4:]) - return n - 4, err + var n int + var callErr error + err = rc.Read(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &to[0], Len: uint64(len(to))}, + } + n, callErr = tunReadv(int(fd), iovecs) + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + if n < 4 { + return 0, nil + } + return n - 4, nil } -// Write is only valid for single threaded use +// Write pushes one IP packet onto the tun device. func (t *tun) Write(from []byte) (int, error) { - buf := t.out - if cap(buf) < len(from)+4 { - buf = make([]byte, len(from)+4) - t.out = buf - } - buf = buf[:len(from)+4] - if len(from) == 0 { return 0, syscall.EIO } - // Determine the IP Family for the NULL L2 Header ipVer := from[0] >> 4 - if ipVer == 4 { - buf[3] = syscall.AF_INET - } else if ipVer == 6 { - buf[3] = syscall.AF_INET6 - } else { + var head [4]byte + switch ipVer { + case 4: + head[3] = syscall.AF_INET + case 6: + head[3] = syscall.AF_INET6 + default: return 0, fmt.Errorf("unable to determine IP version from packet") } - copy(buf[4:], from) + // Grab rc as a local so the compiler can devirtualize the call and keep the closure on the stack. + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - n, err := t.f.Write(buf) - return n - 4, err + var n int + var callErr error + err = rc.Write(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &from[0], Len: uint64(len(from))}, + } + n, callErr = tunWritev(int(fd), iovecs) + // Type-assert to syscall.Errno so the EAGAIN/EWOULDBLOCK/EINTR check doesn't box the errno + // constants into error interfaces on every call. + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + + return n - 4, nil } func (t *tun) addIp(cidr netip.Prefix) error { From 1e66c0d3eefd66fee1ff88dda1ec9d565366ac8f Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 17:32:25 -0500 Subject: [PATCH 15/25] hostmap: unlink a multi-vpnAddr hostinfo from the shared chain exactly once on delete (#1788) --- e2e/handshakes_test.go | 75 ++++++++++++++++++++++++++++++ hostmap.go | 52 ++++++++------------- hostmap_test.go | 101 +++++++++++++++++++++++++++++++++++++++++ 3 files changed, 195 insertions(+), 33 deletions(-) diff --git a/e2e/handshakes_test.go b/e2e/handshakes_test.go index d0b9543c..d580eb21 100644 --- a/e2e/handshakes_test.go +++ b/e2e/handshakes_test.go @@ -1535,3 +1535,78 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) { myControl.Stop() theirControl.Stop() } + +func TestMultiVpnAddrDeletePrimaryKeepsSecondAddr(t *testing.T) { + t.Parallel() + // Regression for the hostmap multi-vpnAddr delete bug. A dual-stack (v4+v6) V2-cert peer that + // handshakes twice at once ends up with two hostinfos linked in the shared next/prev chain, with the + // primary owning both addresses. Deleting that primary (e.g. connection manager dropping it, a + // CloseTunnel, a collision) must promote the surviving sibling for EVERY address. The pre-fix code + // unlinked the chain once per address, so it promoted the sibling for the first address and orphaned + // the second: the peer stayed reachable at its v4 addr but not its v6 addr despite a live tunnel. + + ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) + myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fd00::1/64", nil) + theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24,fd00::2/64", nil) + + // This bug only exists for peers carrying more than one vpn address + require.Len(t, theirVpnIpNet, 2) + theirV4 := theirVpnIpNet[0].Addr() + theirV6 := theirVpnIpNet[1].Addr() + + // Put their info in our lighthouse and vice versa + myControl.InjectLightHouseAddr(theirV4, theirUdpAddr) + theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr) + + // Build a router so we don't have to reason who gets which packet + r := router.NewR(t, myControl, theirControl) + defer r.RenderFlow() + + myControl.Start() + theirControl.Start() + + // Race a handshake so both of us build a hostinfo for the other, leaving my hostmap with a single + // host (them) backed by two linked hostinfos, just like TestStage1Race. + myControl.InjectTunPacket(BuildTunUDPPacket(theirV4, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))) + theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirV4, 80, []byte("Hi from them"))) + + myHsForThem := myControl.GetFromUDP(true) + theirHsForMe := theirControl.GetFromUDP(true) + + r.InjectUDPPacket(theirControl, myControl, theirHsForMe) + r.InjectUDPPacket(myControl, theirControl, myHsForThem) + + r.RouteForAllUntilTxTun(theirControl) + r.RouteForAllUntilTxTun(myControl) + + r.RenderHostmaps("Racing hostmaps", myControl, theirControl) + + // Two hostinfos for them means the shared next/prev chain has a sibling to promote. The Hosts map has + // one entry per vpn address (two, for dual stack), so the index count is what tells us there are two + // hostinfos. + require.Len(t, myControl.ListHostmapIndexes(false), 2) + + // The primary owns both of their addresses + primaryV4 := myControl.GetHostInfoByVpnAddr(theirV4, false) + primaryV6 := myControl.GetHostInfoByVpnAddr(theirV6, false) + require.NotNil(t, primaryV4) + require.NotNil(t, primaryV6) + require.Equal(t, primaryV4.LocalIndex, primaryV6.LocalIndex, "both addrs should point at the same primary") + + // Delete the primary tunnel. localOnly so we don't perturb their side, we only care about my hostmap. + require.True(t, myControl.CloseTunnel(theirV4, true)) + + // The surviving sibling must still serve BOTH addresses. + survivorV4 := myControl.GetHostInfoByVpnAddr(theirV4, false) + survivorV6 := myControl.GetHostInfoByVpnAddr(theirV6, false) + require.NotNil(t, survivorV4, "v4 addr should still resolve to the surviving tunnel") + // Pre-fix this is nil: the second address was orphaned when the primary was deleted. + require.NotNil(t, survivorV6, "v6 addr was orphaned after deleting the primary (multi-vpnAddr delete bug)") + assert.Equal(t, survivorV4.LocalIndex, survivorV6.LocalIndex, "both addrs should promote to the same survivor") + assert.NotEqual(t, primaryV4.LocalIndex, survivorV4.LocalIndex, "a different hostinfo should now be primary") + + r.RenderHostmaps("Final hostmaps", myControl, theirControl) + + myControl.Stop() + theirControl.Stop() +} diff --git a/hostmap.go b/hostmap.go index e7dd17a0..e7399034 100644 --- a/hostmap.go +++ b/hostmap.go @@ -438,43 +438,29 @@ func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) { } func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) { + isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil + for _, addr := range hostinfo.vpnAddrs { - h := hm.Hosts[addr] - for h != nil { - if h == hostinfo { - hm.unlockedInnerDeleteHostInfo(h, addr) - } - h = h.next + if hm.Hosts[addr] != hostinfo { + continue + } + if hostinfo.next != nil { + // Promote the next hostinfo in the shared chain to primary for this address + hm.Hosts[addr] = hostinfo.next + } else { + delete(hm.Hosts, addr) } } -} + if len(hm.Hosts) == 0 { + hm.Hosts = map[netip.Addr]*HostInfo{} + } -func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Addr) { - primary, ok := hm.Hosts[addr] - isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil - if ok && primary == hostinfo { - // The vpn addr pointer points to the same hostinfo as the local index id, we can remove it - delete(hm.Hosts, addr) - if len(hm.Hosts) == 0 { - hm.Hosts = map[netip.Addr]*HostInfo{} - } - - if hostinfo.next != nil { - // We had more than 1 hostinfo at this vpn addr, promote the next in the list to primary - hm.Hosts[addr] = hostinfo.next - // It is primary, there is no previous hostinfo now - hostinfo.next.prev = nil - } - - } else { - // Relink if we were in the middle of multiple hostinfos for this vpn addr - if hostinfo.prev != nil { - hostinfo.prev.next = hostinfo.next - } - - if hostinfo.next != nil { - hostinfo.next.prev = hostinfo.prev - } + // Splice this hostinfo out of the shared chain exactly once + if hostinfo.prev != nil { + hostinfo.prev.next = hostinfo.next + } + if hostinfo.next != nil { + hostinfo.next.prev = hostinfo.prev } hostinfo.next = nil diff --git a/hostmap_test.go b/hostmap_test.go index 2bd7bd43..156444a3 100644 --- a/hostmap_test.go +++ b/hostmap_test.go @@ -194,6 +194,107 @@ func TestHostMap_DeleteHostInfo(t *testing.T) { assert.Nil(t, prim) } +// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one +// vpnAddr and shares its next/prev chain with a live sibling. Deleting the head must not corrupt the +// sibling: every address the sibling owns has to keep pointing at it. The pre-fix code unlinked the shared +// chain once per vpnAddr, so on the first address it nil'd next/prev, and on the second address the node +// looked already-detached: it dropped the map entry instead of promoting the sibling (and tripped the +// isLastHostinfo relay teardown). See unlockedDeleteHostInfo. +func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + + f := &Interface{} + + a := netip.MustParseAddr("0.0.0.1") + b := netip.MustParseAddr("0.0.0.2") + + // Two tunnels for the same peer, each reachable at both a and b. + other := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 1} + head := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2} + + hm.unlockedAddHostInfo(other, f) + hm.unlockedAddHostInfo(head, f) + + // head is primary for both addresses, other is next in the shared chain + assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(a).localIndexId) + assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(b).localIndexId) + assert.Equal(t, other.localIndexId, head.next.localIndexId) + assert.Equal(t, head.localIndexId, other.prev.localIndexId) + + // Delete the head. other is still live, so it must become primary for BOTH addresses. + hm.DeleteHostInfo(head) + + // Pre-fix: QueryVpnAddr(b) came back nil here because the second address was deleted rather than + // promoted, leaving other unreachable at b. + require.NotNil(t, hm.QueryVpnAddr(a)) + require.NotNil(t, hm.QueryVpnAddr(b)) + assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(a).localIndexId) + assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(b).localIndexId) + + // other is now the only hostinfo in the chain + assert.Nil(t, other.prev) + assert.Nil(t, other.next) + + // head is fully detached + assert.Nil(t, head.prev) + assert.Nil(t, head.next) + assert.Nil(t, hm.QueryIndex(head.localIndexId)) +} + +// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune +// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long) +// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is +// primary for none of the addresses, and both address chains must stay consistent afterwards. +func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + + f := &Interface{} + + a := netip.MustParseAddr("0.0.0.1") + b := netip.MustParseAddr("0.0.0.2") + + // Add one more than the cap, newest last so it becomes head. Every hostinfo owns both a and b. + hostinfos := make([]*HostInfo, 0, MaxHostInfosPerVpnIp+1) + for i := 0; i <= MaxHostInfosPerVpnIp; i++ { + hostinfos = append(hostinfos, &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: uint32(i + 1)}) + } + // Add oldest first (highest index in our slice) so the very first one added is the overflow victim. + for i := len(hostinfos) - 1; i >= 0; i-- { + hm.unlockedAddHostInfo(hostinfos[i], f) + } + + oldest := hostinfos[len(hostinfos)-1] + + // The oldest hostinfo should have been pruned and fully detached + assert.Nil(t, oldest.next) + assert.Nil(t, oldest.prev) + assert.Nil(t, hm.QueryIndex(oldest.localIndexId)) + + // Both addresses resolve to the same head, and that head is one of the survivors (not the pruned one) + primA := hm.QueryVpnAddr(a) + primB := hm.QueryVpnAddr(b) + require.NotNil(t, primA) + require.NotNil(t, primB) + assert.Equal(t, primA.localIndexId, primB.localIndexId) + assert.NotEqual(t, oldest.localIndexId, primA.localIndexId) + + // Walk the shared chain: exactly MaxHostInfosPerVpnIp survivors, no cycles, oldest absent + seen := map[uint32]struct{}{} + for h := primA; h != nil; h = h.next { + _, dup := seen[h.localIndexId] + require.False(t, dup, "cycle detected in hostinfo chain") + seen[h.localIndexId] = struct{}{} + if h.next != nil { + assert.Equal(t, h.localIndexId, h.next.prev.localIndexId, "prev pointer must mirror next") + } + } + assert.Len(t, seen, MaxHostInfosPerVpnIp) + _, prunedStillPresent := seen[oldest.localIndexId] + assert.False(t, prunedStillPresent) +} + func TestHostMap_reload(t *testing.T) { l := test.NewLogger() c := config.NewC(test.NewLogger()) From c1eea118f4e8183721d16d2a23db8d33da6ffd88 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Tue, 7 Jul 2026 20:43:26 -0500 Subject: [PATCH 16/25] fix firewall port/proto bypass in parseV6 from uint8 extension-header length overflow (#1789) --- outside.go | 6 ++---- outside_test.go | 35 +++++++++++++++++++++++++++++++++++ 2 files changed, 37 insertions(+), 4 deletions(-) diff --git a/outside.go b/outside.go index 29607eae..4464acdf 100644 --- a/outside.go +++ b/outside.go @@ -422,16 +422,14 @@ func parseV6(data []byte, incoming bool, fp *firewall.Packet) error { if dataLen <= offset+1 { break } - - next = int(data[offset+1]+2) << 2 + next = (int(data[offset+1]) + 2) << 2 default: // Normal ipv6 header length processing if dataLen <= offset+1 { break } - - next = int(data[offset+1]+1) << 3 + next = (int(data[offset+1]) + 1) << 3 } if next <= 0 { diff --git a/outside_test.go b/outside_test.go index 042ccbb3..4a24cae5 100644 --- a/outside_test.go +++ b/outside_test.go @@ -640,3 +640,38 @@ func serializeAH(ah *layers.IPSecAH) []byte { return buf.Bytes() } + +// Test_newPacket_v6ExtHeaderOverflow is a regression test for the IPv6 extension-header +// length uint8 overflow in parseV6. A Destination-Options header with HdrExtLen=255 spans +// (255+1)*8 = 2048 bytes, so the real transport header sits at offset 2088. Before the fix +// the advance was computed in uint8 and wrapped to 0 (then clamped to 8), so the firewall +// read the transport header ~2KB too early from attacker-controlled option bytes while the +// host OS parses the real header, a firewall port/proto bypass. The fix makes parseV6 land +// on the same offset the host does. +func Test_newPacket_v6ExtHeaderOverflow(t *testing.T) { + p := &firewall.Packet{} + + const ( + hdrLen = 40 // IPv6 header + extLen = 2048 // (255+1)*8, the true Destination-Options header size + realTCPAt = hdrLen + extLen // 2088, where the host reads the transport header + forgedTCPAt = hdrLen + 8 // 48, where the pre-fix wrapped+clamped walk landed + ) + + pkt := make([]byte, realTCPAt+4) + pkt[0] = 0x60 // version 6 + pkt[6] = byte(layers.IPProtocolIPv6Destination) // NextHeader -> Destination Options + 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. + binary.BigEndian.PutUint16(pkt[forgedTCPAt+2:forgedTCPAt+4], 443) + // Real transport header at the offset the host actually uses: dst port 22. + binary.BigEndian.PutUint16(pkt[realTCPAt+2:realTCPAt+4], 22) + + require.NoError(t, newPacket(pkt, true, p)) + 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") +} From 384610f81ac7b41b9bd80ce260d8493250436f2a Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Thu, 9 Jul 2026 09:34:30 -0500 Subject: [PATCH 17/25] hostmap: replace the shared next/prev hostinfo chain with independent per-address lists so divergent or overlapping vpnAddr sets cannot corrupt the map (#1790) --- handshake_manager.go | 7 +- hostmap.go | 242 +++++++++++++++---------- hostmap_test.go | 418 ++++++++++++++++++++++++------------------- 3 files changed, 386 insertions(+), 281 deletions(-) diff --git a/handshake_manager.go b/handshake_manager.go index 0d25305f..913918c2 100644 --- a/handshake_manager.go +++ b/handshake_manager.go @@ -430,14 +430,11 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket // Check if we already have a tunnel with this vpn ip existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]] if found && existingHostInfo != nil { - testHostInfo := existingHostInfo - for testHostInfo != nil { - // Is it just a delayed handshake packet? + // Is it just a delayed handshake packet? Check every hostinfo we hold for this address. + for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) { if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) { return testHostInfo, ErrAlreadySeen } - - testHostInfo = testHostInfo.next } // Is this a newer handshake? diff --git a/hostmap.go b/hostmap.go index e7399034..2f2db101 100644 --- a/hostmap.go +++ b/hostmap.go @@ -56,11 +56,20 @@ type Relay struct { } type HostMap struct { - sync.RWMutex //Because we concurrently read and write to our maps - Indexes map[uint32]*HostInfo - Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object - RemoteIndexes map[uint32]*HostInfo + sync.RWMutex //Because we concurrently read and write to our maps + Indexes map[uint32]*HostInfo + Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object + RemoteIndexes map[uint32]*HostInfo + // Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel + // for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores + // the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a]. + // Each address gets its own independent list, so a hostinfo owning multiple addresses can + // never corrupt another address's ordering the way the old shared next/prev chain could. + // Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written + // directly only in the single-hostinfo fast paths where moreHosts is known to have no entry, + // and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains. Hosts map[netip.Addr]*HostInfo + moreHosts map[netip.Addr][]*HostInfo preferredRanges atomic.Pointer[[]netip.Prefix] l *slog.Logger } @@ -266,10 +275,6 @@ type HostInfo struct { lastRoam time.Time lastRoamRemote netip.AddrPort - // Used to track other hostinfos for this vpn ip since only 1 can be primary - // Synchronised via hostmap lock and not the hostinfo lock. - next, prev *HostInfo - //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 @@ -334,6 +339,7 @@ func newHostMap(l *slog.Logger) *HostMap { Relays: map[uint32]*HostInfo{}, RemoteIndexes: map[uint32]*HostInfo{}, Hosts: map[netip.Addr]*HostInfo{}, + moreHosts: map[netip.Addr][]*HostInfo{}, l: l, } } @@ -382,13 +388,55 @@ func (hm *HostMap) EmitStats() { metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen)) } -// DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip +// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty +// list removes the address. This is the one place Hosts and moreHosts are written together, keep +// it that way. Callers must hold the write lock. +func (hm *HostMap) unlockedSetHostsForAddr(addr netip.Addr, list []*HostInfo) { + if len(list) == 0 { + delete(hm.Hosts, addr) + delete(hm.moreHosts, addr) + return + } + hm.Hosts[addr] = list[0] + if len(list) > 1 { + hm.moreHosts[addr] = list + } else { + delete(hm.moreHosts, addr) + } +} + +// unlockedGetHostList returns every hostinfo holding addr, primary first, or nil if we have no +// tunnel for addr. The common single-hostinfo case builds a fresh one element list, so keep this +// off the packet hot path; the primary is a direct Hosts read. Callers must hold the lock (read +// or write). +func (hm *HostMap) unlockedGetHostList(addr netip.Addr) []*HostInfo { + if list, ok := hm.moreHosts[addr]; ok { + return list + } + if h, ok := hm.Hosts[addr]; ok { + return []*HostInfo{h} + } + return nil +} + +// removeHostInfo returns list with hi removed (order preserved), or list unchanged if hi is +// absent. It deletes in place: every mutator holds the hostmap write lock and no reader ever +// retains a slice across a mutation (readers iterate under RLock), so there is no snapshot to +// invalidate. +func removeHostInfo(list []*HostInfo, hi *HostInfo) []*HostInfo { + idx := slices.Index(list, hi) + if idx < 0 { + return list + } + return slices.Delete(list, idx, idx+1) +} + +// DeleteHostInfo will fully unlink the hostinfo and return true if no other hostinfo still holds +// any of its vpn addrs, meaning we no longer have a tunnel to the peer func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool { // Delete the host itself, ensuring it's not modified anymore hm.Lock() - // If we have a previous or next hostinfo then we are not the last one for this vpn ip - final := (hostinfo.next == nil && hostinfo.prev == nil) - hm.unlockedDeleteHostInfo(hostinfo) + final := hm.unlockedDeleteHostInfo(hostinfo) hm.Unlock() return final @@ -401,70 +449,62 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) { } func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) { - // Get the current primary, if it exists - oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]] - - // Every address in the hostinfo gets elevated to primary - for _, vpnAddr := range hostinfo.vpnAddrs { - //NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on - // indexes so it should be fine. - hm.Hosts[vpnAddr] = hostinfo - } - - // If we are already primary then we won't bother re-linking - if oldHostinfo == hostinfo { + // A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race + // tunnel teardown, deciding to promote under the read lock and only taking the write lock + // after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every + // live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test. + if hm.Indexes[hostinfo.localIndexId] != hostinfo { return } - // Unlink this hostinfo - if hostinfo.prev != nil { - hostinfo.prev.next = hostinfo.next - } - if hostinfo.next != nil { - hostinfo.next.prev = hostinfo.prev - } - - // If there wasn't a previous primary then clear out any links - if oldHostinfo == nil { - hostinfo.next = nil - hostinfo.prev = nil - return - } - - // Relink the hostinfo as primary - hostinfo.next = oldHostinfo - oldHostinfo.prev = hostinfo - hostinfo.prev = nil -} - -func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) { - isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil - + // Move hostinfo to the front (primary) of each of its address lists. The lists are + // independent per address, so this can never leave a dangling entry the way promoting + // against a single shared chain could. for _, addr := range hostinfo.vpnAddrs { - if hm.Hosts[addr] != hostinfo { + if hm.Hosts[addr] == hostinfo { + // Already primary for this address, the list is already in the right order continue } - if hostinfo.next != nil { - // Promote the next hostinfo in the shared chain to primary for this address - hm.Hosts[addr] = hostinfo.next - } else { - delete(hm.Hosts, addr) + list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo) + list = append([]*HostInfo{hostinfo}, list...) + hm.unlockedSetHostsForAddr(addr, list) + } +} + +// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index +// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have +// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse +// state and disestablish relays. +func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool { + // Remove this hostinfo from each of its address lists. The lists are independent, so a + // sibling is never promoted to an address it does not own and no other list is touched. + final := true + for _, addr := range hostinfo.vpnAddrs { + if list, ok := hm.moreHosts[addr]; ok { + list = removeHostInfo(list, hostinfo) + hm.unlockedSetHostsForAddr(addr, list) + if len(list) > 0 { + final = false + } + } else if existing, ok := hm.Hosts[addr]; ok { + if existing == hostinfo { + // Common case, the only hostinfo for this address. moreHosts has no entry to clean up. + delete(hm.Hosts, addr) + } else { + // We don't hold this address but another hostinfo does, we still have a tunnel to the peer + final = false + } } } + + // Go maps never shrink their buckets, replace fully drained maps so a node that churned + // through a large peer count gives the memory back. Same idiom as the index maps below. if len(hm.Hosts) == 0 { hm.Hosts = map[netip.Addr]*HostInfo{} } - - // Splice this hostinfo out of the shared chain exactly once - if hostinfo.prev != nil { - hostinfo.prev.next = hostinfo.next + if len(hm.moreHosts) == 0 { + hm.moreHosts = map[netip.Addr][]*HostInfo{} } - if hostinfo.next != nil { - hostinfo.next.prev = hostinfo.prev - } - - hostinfo.next = nil - hostinfo.prev = nil // The remote index uses index ids outside our control so lets make sure we are only removing // the remote index pointer here if it points to the hostinfo we are deleting @@ -488,7 +528,7 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) { ) } - if isLastHostinfo { + if final { // I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next // hops as 'Requested' so that new relay tunnels are created in the future. hm.unlockedDisestablishVpnAddrRelayFor(hostinfo) @@ -497,6 +537,8 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) { for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() { delete(hm.Relays, localRelayIdx) } + + return final } func (hm *HostMap) QueryIndex(index uint32) *HostInfo { @@ -540,19 +582,30 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net hm.RLock() defer hm.RUnlock() + // This runs per relayed packet, so check the primary with a single map probe and only consult + // moreHosts when the primary can't relay for us. h, ok := hm.Hosts[relayHostIp] if !ok { return nil, nil, errors.New("unable to find host") } - for h != nil { - for _, targetIp := range targetIps { - r, ok := h.relayState.QueryRelayForByIp(targetIp) - if ok && r.State == Established { - return h, r, nil + for _, targetIp := range targetIps { + r, ok := h.relayState.QueryRelayForByIp(targetIp) + if ok && r.State == Established { + return h, r, nil + } + } + + if list, ok := hm.moreHosts[relayHostIp]; ok { + // list[0] is the primary we already checked + for _, h := range list[1:] { + for _, targetIp := range targetIps { + r, ok := h.relayState.QueryRelayForByIp(targetIp) + if ok && r.State == Established { + return h, r, nil + } } } - h = h.next } return nil, nil, errors.New("unable to find host with relay") @@ -560,20 +613,14 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) { for _, relayHostIp := range hi.relayState.CopyRelayIps() { - if h, ok := hm.Hosts[relayHostIp]; ok { - for h != nil { - h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished) - h = h.next - } + for _, h := range hm.unlockedGetHostList(relayHostIp) { + h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished) } } for _, rs := range hi.relayState.CopyAllRelayFor() { if rs.Type == ForwardingType { - if h, ok := hm.Hosts[rs.PeerAddr]; ok { - for h != nil { - h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished) - h = h.next - } + for _, h := range hm.unlockedGetHostList(rs.PeerAddr) { + h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished) } } } @@ -623,22 +670,27 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) { } func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) { - existing := hm.Hosts[vpnAddr] - hm.Hosts[vpnAddr] = hostinfo - - if existing != nil && existing != hostinfo { - hostinfo.next = existing - existing.prev = hostinfo + existing, ok := hm.Hosts[vpnAddr] + if !ok { + // Common case, the first hostinfo for this address. moreHosts stays empty. + hm.Hosts[vpnAddr] = hostinfo + return } - i := 1 - check := hostinfo - for check != nil { - if i > MaxHostInfosPerVpnIp { - hm.unlockedDeleteHostInfo(check) - } - check = check.next - i++ + // The new hostinfo becomes the primary for this address. Remove any stale copy of it first so + // we never hold a duplicate, then prepend. + list, ok := hm.moreHosts[vpnAddr] + if !ok { + list = []*HostInfo{existing} + } + list = removeHostInfo(list, hostinfo) + list = append([]*HostInfo{hostinfo}, list...) + hm.unlockedSetHostsForAddr(vpnAddr, list) + + // Enforce the per-address cap by fully retiring the oldest hostinfo once we exceed it. + // Deleting it removes it from all of its addresses and the index maps, matching prior behavior. + if len(list) > MaxHostInfosPerVpnIp { + hm.unlockedDeleteHostInfo(list[len(list)-1]) } } diff --git a/hostmap_test.go b/hostmap_test.go index 156444a3..9cfebe17 100644 --- a/hostmap_test.go +++ b/hostmap_test.go @@ -2,6 +2,7 @@ package nebula import ( "net/netip" + "slices" "testing" "github.com/slackhq/nebula/config" @@ -10,78 +11,84 @@ import ( "github.com/stretchr/testify/require" ) +// chainIds returns the localIndexIds of the hostinfos holding addr, primary (index 0) first. It +// also validates the Hosts/moreHosts sync contract on every call so a mutation that broke it +// fails fast. +func chainIds(t *testing.T, hm *HostMap, addr netip.Addr) []uint32 { + t.Helper() + assertHostMapInvariants(t, hm) + list := hm.unlockedGetHostList(addr) + ids := make([]uint32, len(list)) + for i, h := range list { + ids[i] = h.localIndexId + } + return ids +} + +// assertHostMapInvariants checks the Hosts/moreHosts contract: moreHosts only holds addresses +// with 2 or more hostinfos, its first entry is always the primary in Hosts, lists never hold +// duplicates, every hostinfo in a list owns the address and is registered in Indexes, and every +// indexed hostinfo is reachable through each of its addresses. +func assertHostMapInvariants(t *testing.T, hm *HostMap) { + t.Helper() + for addr, list := range hm.moreHosts { + require.GreaterOrEqualf(t, len(list), 2, "moreHosts[%s] must hold at least 2 hostinfos", addr) + require.Samef(t, hm.Hosts[addr], list[0], "moreHosts[%s][0] must match the primary in Hosts", addr) + seen := map[*HostInfo]bool{} + for _, h := range list { + require.NotNilf(t, h, "moreHosts[%s] must never hold a nil hostinfo", addr) + require.Falsef(t, seen[h], "moreHosts[%s] holds hostinfo %d twice", addr, h.localIndexId) + seen[h] = true + require.Samef(t, hm.Indexes[h.localIndexId], h, "moreHosts[%s] member %d is not registered in Indexes", addr, h.localIndexId) + require.Truef(t, slices.Contains(h.vpnAddrs, addr), "moreHosts[%s] member %d does not own the address", addr, h.localIndexId) + } + } + for addr, h := range hm.Hosts { + require.NotNilf(t, h, "Hosts[%s] must never be nil", addr) + require.Samef(t, hm.Indexes[h.localIndexId], h, "Hosts[%s] primary %d is not registered in Indexes", addr, h.localIndexId) + require.Truef(t, slices.Contains(h.vpnAddrs, addr), "Hosts[%s] primary (index %d) does not own the address", addr, h.localIndexId) + } + for idx, h := range hm.Indexes { + require.Equalf(t, idx, h.localIndexId, "Indexes[%d] holds hostinfo with localIndexId %d", idx, h.localIndexId) + for _, va := range h.vpnAddrs { + require.Truef(t, slices.Contains(hm.unlockedGetHostList(va), h), "indexed hostinfo %d is missing from the list for %s", idx, va) + } + } +} + func TestHostMap_MakePrimary(t *testing.T) { l := test.NewLogger() hm := newHostMap(l) f := &Interface{} + a := netip.MustParseAddr("0.0.0.1") - h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1} - h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2} - h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3} - h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4} + h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} + h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2} + h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3} + h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4} hm.unlockedAddHostInfo(h4, f) hm.unlockedAddHostInfo(h3, f) hm.unlockedAddHostInfo(h2, f) hm.unlockedAddHostInfo(h1, f) - // Make sure we go h1 -> h2 -> h3 -> h4 - prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h1.localIndexId, prim.localIndexId) - assert.Equal(t, h2.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h1.localIndexId, h2.prev.localIndexId) - assert.Equal(t, h3.localIndexId, h2.next.localIndexId) - assert.Equal(t, h2.localIndexId, h3.prev.localIndexId) - assert.Equal(t, h4.localIndexId, h3.next.localIndexId) - assert.Equal(t, h3.localIndexId, h4.prev.localIndexId) - assert.Nil(t, h4.next) + // Most-recently-added is primary: h1, h2, h3, h4 + assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a)) + assert.Equal(t, h1, hm.QueryVpnAddr(a)) - // Swap h3/middle to primary + // Swap the middle to primary: h3, h1, h2, h4 hm.MakePrimary(h3) + assert.Equal(t, []uint32{3, 1, 2, 4}, chainIds(t, hm, a)) + assert.Equal(t, h3, hm.QueryVpnAddr(a)) - // Make sure we go h3 -> h1 -> h2 -> h4 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h3.localIndexId, prim.localIndexId) - assert.Equal(t, h1.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h2.localIndexId, h1.next.localIndexId) - assert.Equal(t, h3.localIndexId, h1.prev.localIndexId) - assert.Equal(t, h4.localIndexId, h2.next.localIndexId) - assert.Equal(t, h1.localIndexId, h2.prev.localIndexId) - assert.Equal(t, h2.localIndexId, h4.prev.localIndexId) - assert.Nil(t, h4.next) - - // Swap h4/tail to primary + // Swap the tail to primary: h4, h3, h1, h2 hm.MakePrimary(h4) + assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a)) - // Make sure we go h4 -> h3 -> h1 -> h2 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h4.localIndexId, prim.localIndexId) - assert.Equal(t, h3.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h1.localIndexId, h3.next.localIndexId) - assert.Equal(t, h4.localIndexId, h3.prev.localIndexId) - assert.Equal(t, h2.localIndexId, h1.next.localIndexId) - assert.Equal(t, h3.localIndexId, h1.prev.localIndexId) - assert.Equal(t, h1.localIndexId, h2.prev.localIndexId) - assert.Nil(t, h2.next) - - // Swap h4 again should be no-op + // Swapping the current primary again is a no-op hm.MakePrimary(h4) - - // Make sure we go h4 -> h3 -> h1 -> h2 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h4.localIndexId, prim.localIndexId) - assert.Equal(t, h3.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h1.localIndexId, h3.next.localIndexId) - assert.Equal(t, h4.localIndexId, h3.prev.localIndexId) - assert.Equal(t, h2.localIndexId, h1.next.localIndexId) - assert.Equal(t, h3.localIndexId, h1.prev.localIndexId) - assert.Equal(t, h1.localIndexId, h2.prev.localIndexId) - assert.Nil(t, h2.next) + assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a)) } func TestHostMap_DeleteHostInfo(t *testing.T) { @@ -89,13 +96,14 @@ func TestHostMap_DeleteHostInfo(t *testing.T) { hm := newHostMap(l) f := &Interface{} + a := netip.MustParseAddr("0.0.0.1") - h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1} - h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2} - h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3} - h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4} - h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5} - h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6} + h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} + h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2} + h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3} + h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4} + h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5} + h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6} hm.unlockedAddHostInfo(h6, f) hm.unlockedAddHostInfo(h5, f) @@ -104,94 +112,110 @@ func TestHostMap_DeleteHostInfo(t *testing.T) { hm.unlockedAddHostInfo(h2, f) hm.unlockedAddHostInfo(h1, f) - // h6 should be deleted - assert.Nil(t, h6.next) - assert.Nil(t, h6.prev) - h := hm.QueryIndex(h6.localIndexId) - assert.Nil(t, h) + // h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first. + assert.Nil(t, hm.QueryIndex(h6.localIndexId)) + assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a)) - // Make sure we go h1 -> h2 -> h3 -> h4 -> h5 - prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h1.localIndexId, prim.localIndexId) - assert.Equal(t, h2.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h1.localIndexId, h2.prev.localIndexId) - assert.Equal(t, h3.localIndexId, h2.next.localIndexId) - assert.Equal(t, h2.localIndexId, h3.prev.localIndexId) - assert.Equal(t, h4.localIndexId, h3.next.localIndexId) - assert.Equal(t, h3.localIndexId, h4.prev.localIndexId) - assert.Equal(t, h5.localIndexId, h4.next.localIndexId) - assert.Equal(t, h4.localIndexId, h5.prev.localIndexId) - assert.Nil(t, h5.next) + // Delete primary; not final since siblings remain. + assert.False(t, hm.DeleteHostInfo(h1)) + assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a)) - // Delete primary - hm.DeleteHostInfo(h1) - assert.Nil(t, h1.prev) - assert.Nil(t, h1.next) + // Deleting the same hostinfo again must not report final while siblings remain and must not + // disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a + // second delete looked final and wiped lighthouse state out from under the live sibling. + assert.False(t, hm.DeleteHostInfo(h1)) + assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a)) - // Make sure we go h2 -> h3 -> h4 -> h5 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h2.localIndexId, prim.localIndexId) - assert.Equal(t, h3.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h3.localIndexId, h2.next.localIndexId) - assert.Equal(t, h2.localIndexId, h3.prev.localIndexId) - assert.Equal(t, h4.localIndexId, h3.next.localIndexId) - assert.Equal(t, h3.localIndexId, h4.prev.localIndexId) - assert.Equal(t, h5.localIndexId, h4.next.localIndexId) - assert.Equal(t, h4.localIndexId, h5.prev.localIndexId) - assert.Nil(t, h5.next) + // Delete a middle node. + assert.False(t, hm.DeleteHostInfo(h3)) + assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a)) - // Delete in the middle - hm.DeleteHostInfo(h3) - assert.Nil(t, h3.prev) - assert.Nil(t, h3.next) + // Delete the tail. + assert.False(t, hm.DeleteHostInfo(h5)) + assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a)) - // Make sure we go h2 -> h4 -> h5 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h2.localIndexId, prim.localIndexId) - assert.Equal(t, h4.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h4.localIndexId, h2.next.localIndexId) - assert.Equal(t, h2.localIndexId, h4.prev.localIndexId) - assert.Equal(t, h5.localIndexId, h4.next.localIndexId) - assert.Equal(t, h4.localIndexId, h5.prev.localIndexId) - assert.Nil(t, h5.next) + // Delete the head; h4 remains and becomes primary. + assert.False(t, hm.DeleteHostInfo(h2)) + assert.Equal(t, []uint32{4}, chainIds(t, hm, a)) + assert.Equal(t, h4, hm.QueryVpnAddr(a)) - // Delete the tail - hm.DeleteHostInfo(h5) - assert.Nil(t, h5.prev) - assert.Nil(t, h5.next) + // Delete the only remaining item; final is true and the address is gone. + assert.True(t, hm.DeleteHostInfo(h4)) + assert.Empty(t, chainIds(t, hm, a)) + assert.Nil(t, hm.QueryVpnAddr(a)) - // Make sure we go h2 -> h4 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h2.localIndexId, prim.localIndexId) - assert.Equal(t, h4.localIndexId, prim.next.localIndexId) - assert.Nil(t, prim.prev) - assert.Equal(t, h4.localIndexId, h2.next.localIndexId) - assert.Equal(t, h2.localIndexId, h4.prev.localIndexId) - assert.Nil(t, h4.next) + // Deleting an already-gone hostinfo is still final; nothing holds the address anymore. + assert.True(t, hm.DeleteHostInfo(h4)) + assert.Empty(t, chainIds(t, hm, a)) +} - // Delete the head - hm.DeleteHostInfo(h2) - assert.Nil(t, h2.prev) - assert.Nil(t, h2.next) +// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with +// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and +// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a +// no-op, not a resurrection that installs an unmanaged primary. +func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + f := &Interface{} + a := netip.MustParseAddr("0.0.0.1") - // Make sure we only have h4 - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Equal(t, h4.localIndexId, prim.localIndexId) - assert.Nil(t, prim.prev) - assert.Nil(t, prim.next) - assert.Nil(t, h4.next) + h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} + h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2} + hm.unlockedAddHostInfo(h1, f) + hm.unlockedAddHostInfo(h2, f) - // Delete the only item - hm.DeleteHostInfo(h4) - assert.Nil(t, h4.prev) - assert.Nil(t, h4.next) + // h1 is fully deleted while another goroutine still holds a pointer to it. + assert.False(t, hm.DeleteHostInfo(h1)) + assert.Equal(t, []uint32{2}, chainIds(t, hm, a)) - // Make sure we have nil - prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1")) - assert.Nil(t, prim) + // The stale promote must not bring it back. + hm.MakePrimary(h1) + assert.Equal(t, []uint32{2}, chainIds(t, hm, a)) + assert.Equal(t, h2, hm.QueryVpnAddr(a)) + assert.Nil(t, hm.QueryIndex(h1.localIndexId)) +} + +// TestHostMap_QueryVpnAddrsRelayFor_NonPrimary makes sure a relay established on an older +// hostinfo is still found after a newer tunnel without relay state takes primary for the same +// address. The lookup checks the primary first and falls back to the rest of the list. +func TestHostMap_QueryVpnAddrsRelayFor_NonPrimary(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + f := &Interface{} + relayAddr := netip.MustParseAddr("0.0.0.9") + target := netip.MustParseAddr("0.0.0.1") + + older := &HostInfo{ + vpnAddrs: []netip.Addr{relayAddr}, + localIndexId: 1, + relayState: RelayState{ + relayForByAddr: map[netip.Addr]*Relay{}, + relayForByIdx: map[uint32]*Relay{}, + }, + } + older.relayState.InsertRelay(target, 100, &Relay{Type: ForwardingType, State: Established, LocalIndex: 100, PeerAddr: target}) + hm.unlockedAddHostInfo(older, f) + + // The relay is found on the primary. + h, r, err := hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr) + require.NoError(t, err) + assert.Equal(t, older, h) + assert.Equal(t, uint32(100), r.LocalIndex) + + // A re-handshake with no relay state takes primary; the established relay on the older + // hostinfo must still be found through the fallback. + newer := &HostInfo{vpnAddrs: []netip.Addr{relayAddr}, localIndexId: 2} + hm.unlockedAddHostInfo(newer, f) + assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, relayAddr)) + + h, r, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr) + require.NoError(t, err) + assert.Equal(t, older, h) + assert.Equal(t, uint32(100), r.LocalIndex) + + // No hostinfo at all is a plain miss. + _, _, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, netip.MustParseAddr("0.0.0.42")) + require.Error(t, err) } // TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one @@ -216,32 +240,82 @@ func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) { hm.unlockedAddHostInfo(other, f) hm.unlockedAddHostInfo(head, f) - // head is primary for both addresses, other is next in the shared chain - assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(a).localIndexId) - assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(b).localIndexId) - assert.Equal(t, other.localIndexId, head.next.localIndexId) - assert.Equal(t, head.localIndexId, other.prev.localIndexId) + // head is primary for both addresses, other is next in each address's list. + assert.Equal(t, head, hm.QueryVpnAddr(a)) + assert.Equal(t, head, hm.QueryVpnAddr(b)) + assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a)) + assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, b)) // Delete the head. other is still live, so it must become primary for BOTH addresses. - hm.DeleteHostInfo(head) + assert.False(t, hm.DeleteHostInfo(head)) + assert.Equal(t, other, hm.QueryVpnAddr(a)) + assert.Equal(t, other, hm.QueryVpnAddr(b)) + assert.Equal(t, []uint32{1}, chainIds(t, hm, a)) + assert.Equal(t, []uint32{1}, chainIds(t, hm, b)) - // Pre-fix: QueryVpnAddr(b) came back nil here because the second address was deleted rather than - // promoted, leaving other unreachable at b. - require.NotNil(t, hm.QueryVpnAddr(a)) - require.NotNil(t, hm.QueryVpnAddr(b)) - assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(a).localIndexId) - assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(b).localIndexId) - - // other is now the only hostinfo in the chain - assert.Nil(t, other.prev) - assert.Nil(t, other.next) - - // head is fully detached - assert.Nil(t, head.prev) - assert.Nil(t, head.next) + // head is fully removed from the index map. assert.Nil(t, hm.QueryIndex(head.localIndexId)) } +// TestHostMap_DeleteHostInfo_DivergentVpnAddrs covers chained hostinfos for the same peer whose +// vpnAddrs sets differ (a re-handshake cert added a second address). Deleting the superset node +// must not promote a sibling to an address it does not own. +func TestHostMap_DeleteHostInfo_DivergentVpnAddrs(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + f := &Interface{} + a := netip.MustParseAddr("0.0.0.1") + b := netip.MustParseAddr("0.0.0.2") + + // sub owns only a; super (a newer handshake) owns a and b. + sub := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} + super := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2} + hm.unlockedAddHostInfo(sub, f) + hm.unlockedAddHostInfo(super, f) + + assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a)) + assert.Equal(t, []uint32{2}, chainIds(t, hm, b)) + + // Delete super: a promotes to sub (which owns it); b has no remaining owner and must be + // removed, not dangled at sub (which does not own b). + assert.False(t, hm.DeleteHostInfo(super)) + assert.Equal(t, []uint32{1}, chainIds(t, hm, a)) + assert.Empty(t, chainIds(t, hm, b)) + assert.Equal(t, sub, hm.QueryVpnAddr(a)) + assert.Nil(t, hm.QueryVpnAddr(b)) + assert.Nil(t, hm.QueryIndex(super.localIndexId)) + + // Deleting sub cleans up fully. + assert.True(t, hm.DeleteHostInfo(sub)) + assert.Nil(t, hm.QueryVpnAddr(a)) + assertHostMapInvariants(t, hm) +} + +// TestHostMap_AddDivergentOverlap covers a new hostinfo claiming addresses currently owned by two +// DIFFERENT hostinfos. The old single shared next/prev chain overwrote a pointer and orphaned one +// of them (in Indexes but unreachable via its address); independent per-address lists cannot. +func TestHostMap_AddDivergentOverlap(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + f := &Interface{} + a := netip.MustParseAddr("0.0.0.1") + b := netip.MustParseAddr("0.0.0.2") + + hiA := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} + hiP := &HostInfo{vpnAddrs: []netip.Addr{b}, localIndexId: 2} + hm.unlockedAddHostInfo(hiA, f) + hm.unlockedAddHostInfo(hiP, f) + + hiB := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 3} + hm.unlockedAddHostInfo(hiB, f) + + assert.Equal(t, []uint32{3, 1}, chainIds(t, hm, a)) + assert.Equal(t, []uint32{3, 2}, chainIds(t, hm, b)) + // hiA is still reachable via its address (not orphaned) and still indexed. + assert.Contains(t, chainIds(t, hm, a), hiA.localIndexId) + assert.NotNil(t, hm.QueryIndex(hiA.localIndexId)) +} + // TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune // (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long) // still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is @@ -267,32 +341,14 @@ func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) { oldest := hostinfos[len(hostinfos)-1] - // The oldest hostinfo should have been pruned and fully detached - assert.Nil(t, oldest.next) - assert.Nil(t, oldest.prev) + // The oldest hostinfo was pruned from both lists and the index map. assert.Nil(t, hm.QueryIndex(oldest.localIndexId)) - // Both addresses resolve to the same head, and that head is one of the survivors (not the pruned one) - primA := hm.QueryVpnAddr(a) - primB := hm.QueryVpnAddr(b) - require.NotNil(t, primA) - require.NotNil(t, primB) - assert.Equal(t, primA.localIndexId, primB.localIndexId) - assert.NotEqual(t, oldest.localIndexId, primA.localIndexId) - - // Walk the shared chain: exactly MaxHostInfosPerVpnIp survivors, no cycles, oldest absent - seen := map[uint32]struct{}{} - for h := primA; h != nil; h = h.next { - _, dup := seen[h.localIndexId] - require.False(t, dup, "cycle detected in hostinfo chain") - seen[h.localIndexId] = struct{}{} - if h.next != nil { - assert.Equal(t, h.localIndexId, h.next.prev.localIndexId, "prev pointer must mirror next") - } - } - assert.Len(t, seen, MaxHostInfosPerVpnIp) - _, prunedStillPresent := seen[oldest.localIndexId] - assert.False(t, prunedStillPresent) + // Both addresses hold exactly MaxHostInfosPerVpnIp survivors in the same order; oldest is absent. + require.Len(t, chainIds(t, hm, a), MaxHostInfosPerVpnIp) + assert.Equal(t, chainIds(t, hm, a), chainIds(t, hm, b), "both addresses must list the same survivors in the same order") + assert.NotContains(t, chainIds(t, hm, a), oldest.localIndexId) + assert.Equal(t, hm.QueryVpnAddr(a), hm.QueryVpnAddr(b)) } func TestHostMap_reload(t *testing.T) { From 1b84bd00500a14e8311065d2be61327d873cda03 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Thu, 9 Jul 2026 11:04:25 -0500 Subject: [PATCH 18/25] Remove dev fmt.Println (#1793) --- overlay/tun_freebsd.go | 1 - 1 file changed, 1 deletion(-) diff --git a/overlay/tun_freebsd.go b/overlay/tun_freebsd.go index 3d995553..79f55697 100644 --- a/overlay/tun_freebsd.go +++ b/overlay/tun_freebsd.go @@ -659,7 +659,6 @@ func addRoute(prefix netip.Prefix, gateway netroute.Addr) error { return fmt.Errorf("failed to create route.RouteMessage for change: %w", err) } _, err = unix.Write(sock, data[:]) - fmt.Println("DOING CHANGE") return err } return fmt.Errorf("failed to write route.RouteMessage to socket: %w", err) From 5ecdd4eaa90db965fa338cfa005e2d0a671c80a0 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Thu, 9 Jul 2026 18:48:14 -0500 Subject: [PATCH 19/25] Fix e2e test races when looking at hostmap counts (#1795) --- control_tester.go | 8 ++++++ e2e/handshakes_test.go | 64 ++++++++++++++++++++++-------------------- e2e/tunnels_test.go | 12 ++++---- 3 files changed, 48 insertions(+), 36 deletions(-) diff --git a/control_tester.go b/control_tester.go index 728ac649..422d86ec 100644 --- a/control_tester.go +++ b/control_tester.go @@ -125,6 +125,14 @@ func (c *Control) GetHostmap() *HostMap { return c.f.hostMap } +// GetHostmapIndexCount returns the number of entries in the main hostmap Indexes table, holding +// the hostmap read lock so tests can poll it while connection manager churns tunnels. +func (c *Control) GetHostmapIndexCount() int { + c.f.hostMap.RLock() + defer c.f.hostMap.RUnlock() + return len(c.f.hostMap.Indexes) +} + func (c *Control) GetF() *Interface { return c.f } diff --git a/e2e/handshakes_test.go b/e2e/handshakes_test.go index d580eb21..eb5359cb 100644 --- a/e2e/handshakes_test.go +++ b/e2e/handshakes_test.go @@ -405,7 +405,7 @@ func TestStage1Race(t *testing.T) { r.Log("Spin until connection manager tears down a tunnel") - for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 { + for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) t.Log("Connection manager hasn't ticked yet") time.Sleep(time.Second) @@ -453,9 +453,11 @@ func TestUncleanShutdownRaceLoser(t *testing.T) { r.Log("Nuke my hostmap") myHostmap := myControl.GetHostmap() + myHostmap.Lock() myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{} myHostmap.Indexes = map[uint32]*nebula.HostInfo{} myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{} + myHostmap.Unlock() myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again"))) p = r.RouteForAllUntilTxTun(theirControl) @@ -465,10 +467,10 @@ func TestUncleanShutdownRaceLoser(t *testing.T) { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) r.Log("Wait for the dead index to go away") - start := len(theirControl.GetHostmap().Indexes) + start := theirControl.GetHostmapIndexCount() for { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) - if len(theirControl.GetHostmap().Indexes) < start { + if theirControl.GetHostmapIndexCount() < start { break } time.Sleep(time.Second) @@ -504,9 +506,11 @@ func TestUncleanShutdownRaceWinner(t *testing.T) { r.Log("Nuke my hostmap") theirHostmap := theirControl.GetHostmap() + theirHostmap.Lock() theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{} theirHostmap.Indexes = map[uint32]*nebula.HostInfo{} theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{} + theirHostmap.Unlock() theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again"))) p = r.RouteForAllUntilTxTun(myControl) @@ -517,10 +521,10 @@ func TestUncleanShutdownRaceWinner(t *testing.T) { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) r.Log("Wait for the dead index to go away") - start := len(myControl.GetHostmap().Indexes) + start := myControl.GetHostmapIndexCount() for { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) - if len(myControl.GetHostmap().Indexes) < start { + if myControl.GetHostmapIndexCount() < start { break } time.Sleep(time.Second) @@ -628,10 +632,10 @@ func TestReestablishRelays(t *testing.T) { r.Log("Close the tunnel") relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true) - start := len(myControl.GetHostmap().Indexes) - curIndexes := len(myControl.GetHostmap().Indexes) + start := myControl.GetHostmapIndexCount() + curIndexes := myControl.GetHostmapIndexCount() for curIndexes >= start { - curIndexes = len(myControl.GetHostmap().Indexes) + curIndexes = myControl.GetHostmapIndexCount() r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes) myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail"))) @@ -819,18 +823,18 @@ func TestStage1RaceRelays2(t *testing.T) { t.Log("Wait until we remove extra tunnels") t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d", - len(myControl.GetHostmap().Indexes), - len(theirControl.GetHostmap().Indexes), - len(relayControl.GetHostmap().Indexes), + myControl.GetHostmapIndexCount(), + theirControl.GetHostmapIndexCount(), + relayControl.GetHostmapIndexCount(), ) - hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes) + hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount() retries := 60 for hostInfos > 6 && retries > 0 { - hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes) + hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount() t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d", - len(myControl.GetHostmap().Indexes), - len(theirControl.GetHostmap().Indexes), - len(relayControl.GetHostmap().Indexes), + myControl.GetHostmapIndexCount(), + theirControl.GetHostmapIndexCount(), + relayControl.GetHostmapIndexCount(), ) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) t.Log("Connection manager hasn't ticked yet") @@ -924,24 +928,24 @@ func TestRehandshakingRelays(t *testing.T) { assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl) // We should have two hostinfos on all sides - for len(myControl.GetHostmap().Indexes) != 2 { - t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes)) + for myControl.GetHostmapIndexCount() != 2 { + t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount()) r.Log("Assert the relay tunnel still works") assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.Log("yupitdoes") time.Sleep(time.Second) } t.Logf("myControl hostinfos got cleaned up!") - for len(theirControl.GetHostmap().Indexes) != 2 { - t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes)) + for theirControl.GetHostmapIndexCount() != 2 { + t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount()) r.Log("Assert the relay tunnel still works") assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.Log("yupitdoes") time.Sleep(time.Second) } t.Logf("theirControl hostinfos got cleaned up!") - for len(relayControl.GetHostmap().Indexes) != 2 { - t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes)) + for relayControl.GetHostmapIndexCount() != 2 { + t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount()) r.Log("Assert the relay tunnel still works") assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.Log("yupitdoes") @@ -1029,24 +1033,24 @@ func TestRehandshakingRelaysPrimary(t *testing.T) { assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl) // We should have two hostinfos on all sides - for len(myControl.GetHostmap().Indexes) != 2 { - t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes)) + for myControl.GetHostmapIndexCount() != 2 { + t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount()) r.Log("Assert the relay tunnel still works") assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.Log("yupitdoes") time.Sleep(time.Second) } t.Logf("myControl hostinfos got cleaned up!") - for len(theirControl.GetHostmap().Indexes) != 2 { - t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes)) + for theirControl.GetHostmapIndexCount() != 2 { + t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount()) r.Log("Assert the relay tunnel still works") assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.Log("yupitdoes") time.Sleep(time.Second) } t.Logf("theirControl hostinfos got cleaned up!") - for len(relayControl.GetHostmap().Indexes) != 2 { - t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes)) + for relayControl.GetHostmapIndexCount() != 2 { + t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount()) r.Log("Assert the relay tunnel still works") assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) r.Log("yupitdoes") @@ -1123,7 +1127,7 @@ func TestRehandshaking(t *testing.T) { theirConfig.ReloadConfigString(string(rc)) r.Log("Spin until there is only 1 tunnel") - for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 { + for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) t.Log("Connection manager hasn't ticked yet") time.Sleep(time.Second) @@ -1223,7 +1227,7 @@ func TestRehandshakingLoser(t *testing.T) { myConfig.ReloadConfigString(string(rc)) r.Log("Spin until there is only 1 tunnel") - for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 { + for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 { assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) t.Log("Connection manager hasn't ticked yet") time.Sleep(time.Second) diff --git a/e2e/tunnels_test.go b/e2e/tunnels_test.go index 18c69a3f..7874cc79 100644 --- a/e2e/tunnels_test.go +++ b/e2e/tunnels_test.go @@ -43,8 +43,8 @@ func TestDropInactiveTunnels(t *testing.T) { r.Log("Go inactive and wait for the tunnels to get dropped") waitStart := time.Now() for { - myIndexes := len(myControl.GetHostmap().Indexes) - theirIndexes := len(theirControl.GetHostmap().Indexes) + myIndexes := myControl.GetHostmapIndexCount() + theirIndexes := theirControl.GetHostmapIndexCount() if myIndexes == 0 && theirIndexes == 0 { break } @@ -493,8 +493,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) { waitStart := time.Now() for { - myIndexes := len(myControl.GetHostmap().Indexes) - theirIndexes := len(theirControl.GetHostmap().Indexes) + myIndexes := myControl.GetHostmapIndexCount() + theirIndexes := theirControl.GetHostmapIndexCount() if myIndexes == 0 && theirIndexes == 0 { break } @@ -548,8 +548,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) { r.Log("Injected bogus close tunnel. Let's see!") waitStart = time.Now() for { - myIndexes := len(myControl.GetHostmap().Indexes) - theirIndexes := len(theirControl.GetHostmap().Indexes) + myIndexes := myControl.GetHostmapIndexCount() + theirIndexes := theirControl.GetHostmapIndexCount() if myIndexes == 0 { t.Fatal("myIndexes should not be 0") } From ab736e4c6b34616811ef3344c80bc488ea0b6303 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Fri, 10 Jul 2026 10:35:17 -0500 Subject: [PATCH 20/25] Make Control safe to stop and wait on from any lifecycle state (#1794) --- cmd/nebula-service/main.go | 12 +- cmd/nebula-service/service.go | 31 +++- cmd/nebula/main.go | 5 +- control.go | 72 ++++++--- control_lifecycle_test.go | 292 ++++++++++++++++++++++++++++++++++ interface.go | 43 +++-- main.go | 11 ++ overlay/tun_android.go | 1 + overlay/tun_ios.go | 8 + service/service.go | 24 ++- 10 files changed, 441 insertions(+), 58 deletions(-) create mode 100644 control_lifecycle_test.go diff --git a/cmd/nebula-service/main.go b/cmd/nebula-service/main.go index 724c0c6a..e0b335f5 100644 --- a/cmd/nebula-service/main.go +++ b/cmd/nebula-service/main.go @@ -53,7 +53,12 @@ func main() { l := logging.NewLogger(os.Stdout) if *serviceFlag != "" { - if err := doService(configPath, configTest, Build, serviceFlag); err != nil { + if *configTest { + fmt.Println("-test is not supported with -service, run the config test without -service") + os.Exit(1) + } + + if err := doService(configPath, Build, serviceFlag); err != nil { l.Error("Service command failed", "error", err) os.Exit(1) } @@ -93,15 +98,14 @@ func main() { } if !*configTest { - wait, err := ctrl.Start() - if err != nil { + if err := ctrl.Start(); err != nil { util.LogWithContextIfNeeded("Error while running", err, l) os.Exit(1) } go ctrl.ShutdownBlock() - if err := wait(); err != nil { + if err := ctrl.Wait(); err != nil { l.Error("Nebula stopped due to fatal error", "error", err) os.Exit(2) } diff --git a/cmd/nebula-service/service.go b/cmd/nebula-service/service.go index 7c2b39c8..abe9abe0 100644 --- a/cmd/nebula-service/service.go +++ b/cmd/nebula-service/service.go @@ -3,6 +3,7 @@ package main import ( "fmt" "log" + "os" "github.com/kardianos/service" "github.com/slackhq/nebula" @@ -14,7 +15,6 @@ var logger service.Logger type program struct { configPath *string - configTest *bool build string control *nebula.Control } @@ -40,22 +40,41 @@ func (p *program) Start(s service.Service) error { } }) - p.control, err = nebula.Main(c, *p.configTest, Build, l, nil) + p.control, err = nebula.Main(c, false, Build, l, nil) if err != nil { return err } - p.control.Start() + if err := p.control.Start(); err != nil { + return err + } + + // Nebula can stop itself on a fatal packet reader error, make sure to log it if it happens. + go func() { + if err := p.control.Wait(); err != nil { + logger.Error(fmt.Sprintf("Nebula stopped due to fatal error: %v", err)) + os.Exit(2) + } + }() + return nil } func (p *program) Stop(s service.Service) error { logger.Info("Nebula service stopping.") + if p.control == nil { + return nil + } + p.control.Stop() + + // block until nebula has fully drained before reporting stopped. + // error logging is handled by Start. + _ = p.control.Wait() return nil } -func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error { +func doService(configPath *string, build string, serviceFlag *string) error { if *configPath == "" { p, err := config.DefaultPath() if err != nil { @@ -73,7 +92,6 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag * prg := &program{ configPath: configPath, - configTest: configTest, build: build, } @@ -105,8 +123,9 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag * switch *serviceFlag { case "run": if err := s.Run(); err != nil { - // Route any errors to the system logger + // Route any errors to the system logger and report the failure logger.Error(err) + return err } default: if err := service.Control(s, *serviceFlag); err != nil { diff --git a/cmd/nebula/main.go b/cmd/nebula/main.go index 219519c2..3c786b84 100644 --- a/cmd/nebula/main.go +++ b/cmd/nebula/main.go @@ -84,8 +84,7 @@ func main() { } if !*configTest { - wait, err := ctrl.Start() - if err != nil { + if err := ctrl.Start(); err != nil { util.LogWithContextIfNeeded("Error while running", err, l) os.Exit(1) } @@ -93,7 +92,7 @@ func main() { go ctrl.ShutdownBlock() notifyReady(l) - if err := wait(); err != nil { + if err := ctrl.Wait(); err != nil { l.Error("Nebula stopped due to fatal error", "error", err) os.Exit(2) } diff --git a/control.go b/control.go index 053feab5..a79ebbfa 100644 --- a/control.go +++ b/control.go @@ -69,29 +69,29 @@ type ControlHostInfo struct { } // Start actually runs nebula, this is a nonblocking call. -// The returned function blocks until nebula has fully stopped and returns the -// first fatal reader error (if any). A nil error means nebula shut down -// gracefully; a non-nil error means a reader hit an unexpected failure that -// triggered the shutdown. -func (c *Control) Start() (func() error, error) { +// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown. +func (c *Control) Start() error { c.stateLock.Lock() defer c.stateLock.Unlock() switch c.state { case StateReady: //yay! case StateStopped, StateStopping: - return nil, ErrAlreadyStopped + return ErrAlreadyStopped case StateStarted: - return nil, ErrAlreadyStarted + return ErrAlreadyStarted default: - return nil, ErrUnknownState + return ErrUnknownState } // Activate the interface err := c.f.activate() if err != nil { + // Cancel before Close so a caller returning from Wait always observes a dead Context + c.cancel() + _ = c.f.Close() c.state = StateStopped - return nil, err + return err } // Call all the delayed funcs that waited patiently for the interface to be created. @@ -114,13 +114,9 @@ func (c *Control) Start() (func() error, error) { c.f.triggerShutdown = c.Stop // Start reading packets. - out, err := c.f.run() - if err != nil { - c.state = StateStopped - return nil, err - } + c.f.run() c.state = StateStarted - return out, nil + return nil } func (c *Control) State() RunState { @@ -133,10 +129,26 @@ func (c *Control) Context() context.Context { return c.ctx } -// Stop is a non-blocking call that signals nebula to close all tunnels and shut down +// Stop tears nebula down, closing all tunnels and releasing everything it holds. +// Use Wait to block until the shutdown has completed. +// A Control that has been stopped cannot be started again, Start will return ErrAlreadyStopped. func (c *Control) Stop() { c.stateLock.Lock() - if c.state != StateStarted { + switch c.state { + case StateStarted: + // Fall through to the full teardown below + + case StateReady: + // Never started + c.cancel() + c.state = StateStopped + if err := c.f.Close(); err != nil { + c.l.Error("Close interface failed", "error", err) + } + c.stateLock.Unlock() + return + + default: c.stateLock.Unlock() // We are stopping or stopped already return @@ -145,19 +157,26 @@ func (c *Control) Stop() { c.state = StateStopping c.stateLock.Unlock() - // Stop the handshakeManager (and other services), to prevent new tunnels from - // being created while we're shutting them all down. + // Closing tunnels can be slow with a large hostmap, don't hold the lock for it c.cancel() - c.CloseAllTunnels(false) + + c.stateLock.Lock() + c.state = StateStopped if err := c.f.Close(); err != nil { c.l.Error("Close interface failed", "error", err) } - c.stateLock.Lock() - c.state = StateStopped c.stateLock.Unlock() } +// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error, +// and returns the first fatal packet reader error if there was one. +// It is safe to call from multiple goroutines and at any point in the lifecycle, +// but a Wait on a Control that is never started and never stopped will block forever. +func (c *Control) Wait() error { + return c.f.wait() +} + // ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled func (c *Control) ShutdownBlock() { sigChan := make(chan os.Signal, 1) @@ -170,8 +189,15 @@ func (c *Control) ShutdownBlock() { c.Stop() } -// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change +// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change. func (c *Control) RebindUDPServer() { + c.stateLock.Lock() + defer c.stateLock.Unlock() + + if c.state != StateStarted { + return + } + _ = c.f.outside.Rebind() // Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0 diff --git a/control_lifecycle_test.go b/control_lifecycle_test.go new file mode 100644 index 00000000..0b5d106d --- /dev/null +++ b/control_lifecycle_test.go @@ -0,0 +1,292 @@ +package nebula + +import ( + "context" + "errors" + "io" + "net/netip" + "sync" + "testing" + "time" + + "github.com/gaissmai/bart" + "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/routing" + "github.com/slackhq/nebula/test" + "github.com/slackhq/nebula/udp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type fakeDevice struct { + closeOnce sync.Once + closedCh chan struct{} + closed bool +} + +func newFakeDevice() *fakeDevice { + return &fakeDevice{closedCh: make(chan struct{})} +} + +// Read blocks until Close like a real tun with no traffic, then reports EOF +// the same way a closed device does +func (d *fakeDevice) Read(p []byte) (int, error) { + <-d.closedCh + return 0, io.EOF +} + +func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil } + +func (d *fakeDevice) Close() error { + d.closeOnce.Do(func() { + d.closed = true + close(d.closedCh) + }) + return nil +} + +func (d *fakeDevice) Activate() error { return nil } +func (d *fakeDevice) Networks() []netip.Prefix { return nil } +func (d *fakeDevice) Name() string { return "fake" } +func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil } +func (d *fakeDevice) SupportsMultiqueue() bool { return false } +func (d *fakeDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) { + return nil, errors.New("unsupported") +} + +// newReadyControl hand-builds the minimum Control that Main would have +// produced right before Start, including the construction token NewInterface +// takes so waiters block until Close releases the resources +func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) { + l := test.NewLogger() + dev := newFakeDevice() + conn := &fakeConn{} + ctx, cancel := context.WithCancel(context.Background()) + + myVpnNet := netip.MustParsePrefix("10.128.0.1/16") + nt := new(bart.Lite) + nt.Insert(myVpnNet) + cs := &CertState{ + myVpnNetworks: []netip.Prefix{myVpnNet}, + myVpnNetworksTable: nt, + } + lh, err := NewLightHouseFromConfig(ctx, l, config.NewC(l), cs, nil, nil) + require.NoError(t, err) + + f := &Interface{ + ctx: ctx, + inside: dev, + outside: conn, + writers: []udp.Conn{conn}, + readers: make([]io.ReadWriteCloser, 1), + routines: 1, + hostMap: newHostMap(l), + lightHouse: lh, + l: l, + } + f.wg.Add(1) + + return &Control{ + state: StateReady, + f: f, + l: l, + ctx: ctx, + cancel: cancel, + }, dev, conn +} + +func TestControl_StopBeforeStart(t *testing.T) { + c, dev, conn := newReadyControl(t) + + // A Stop on a never started control must release everything Main acquired + c.Stop() + assert.Equal(t, StateStopped, c.State()) + assert.True(t, dev.closed, "the tun device should have been closed") + assert.True(t, conn.closed, "the udp socket should have been closed") + require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled") + + // Wait must return promptly now that the resources are released + require.NoError(t, c.Wait()) + + // A stopped control can never be started + require.ErrorIs(t, c.Start(), ErrAlreadyStopped) + + // A second Stop is a harmless no-op + c.Stop() + assert.Equal(t, StateStopped, c.State()) + require.NoError(t, c.Wait()) +} + +func TestControl_WaitBlocksUntilStop(t *testing.T) { + c, _, _ := newReadyControl(t) + + done := make(chan error, 1) + go func() { done <- c.Wait() }() + + select { + case <-done: + t.Fatal("Wait returned before Stop") + case <-time.After(50 * time.Millisecond): + } + + c.Stop() + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(time.Second): + t.Fatal("Wait did not return after Stop") + } +} + +type fakeConn struct { + closed bool + rebinds int +} + +func (c *fakeConn) Rebind() error { c.rebinds++; return nil } +func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil } +func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil } +func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil } +func (c *fakeConn) ReloadConfig(_ *config.C) {} +func (c *fakeConn) SupportsMultipleReaders() bool { return true } +func (c *fakeConn) Close() error { c.closed = true; return nil } + +type multiqueueDevice struct { + *fakeDevice +} + +func (d *multiqueueDevice) SupportsMultiqueue() bool { return true } + +func TestControl_StartMultiqueueFailureReleases(t *testing.T) { + dev := &multiqueueDevice{fakeDevice: newFakeDevice()} + conn := &fakeConn{} + ctx, cancel := context.WithCancel(context.Background()) + f := &Interface{ + ctx: ctx, + inside: dev, + outside: conn, + writers: []udp.Conn{conn}, + readers: make([]io.ReadWriteCloser, 2), + routines: 2, + l: test.NewLogger(), + } + f.wg.Add(1) + + c := &Control{ + state: StateReady, + f: f, + l: test.NewLogger(), + ctx: ctx, + cancel: cancel, + } + + // The second reader fails to open, everything must be released + require.Error(t, c.Start()) + assert.Equal(t, StateStopped, c.State()) + assert.True(t, dev.closed, "the tun device should have been closed") + assert.True(t, conn.closed, "the udp socket should have been closed") + require.ErrorIs(t, c.ctx.Err(), context.Canceled) + + // And Wait must not hang on the construction token + require.NoError(t, c.Wait()) +} + +func TestInterface_CloseIsIdempotent(t *testing.T) { + dev := newFakeDevice() + f := &Interface{ + inside: dev, + l: test.NewLogger(), + } + f.wg.Add(1) + + require.NoError(t, f.Close()) + assert.True(t, dev.closed) + + // A second Close must not double release the wg token or the device + require.NoError(t, f.Close()) + require.NoError(t, f.wait()) +} + +func TestControl_FatalErrorReportsThroughWait(t *testing.T) { + c, dev, conn := newReadyControl(t) + + // Mirror what Start wires up, without needing real packet readers + c.f.triggerShutdown = c.Stop + c.state = StateStarted + + boom := errors.New("boom") + c.f.onFatal(boom) + + require.ErrorIs(t, c.Wait(), boom) + assert.Equal(t, StateStopped, c.State()) + assert.True(t, dev.closed) + assert.True(t, conn.closed) + + // A second fatal error must not fire the shutdown again or replace the first + c.f.onFatal(errors.New("later")) + require.ErrorIs(t, c.Wait(), boom) + + // Wait stays factual, a Stop after the death does not mask the error + c.Stop() + require.ErrorIs(t, c.Wait(), boom) +} + +func TestControl_ConcurrentStopAndStart(t *testing.T) { + c, _, _ := newReadyControl(t) + + var wg sync.WaitGroup + for i := 0; i < 2; i++ { + wg.Go(func() { c.Stop() }) + } + wg.Go(func() { _ = c.Start() }) + wg.Go(func() { + _ = c.Wait() + // A returned Wait must always observe the final state, no matter how + // the race resolved + assert.Equal(t, StateStopped, c.State()) + }) + wg.Wait() + + // However the race resolves, the control must end fully stopped with no + // panic and Wait must observe the final state + require.NoError(t, c.Wait()) + assert.Equal(t, StateStopped, c.State()) + require.ErrorIs(t, c.Start(), ErrAlreadyStopped) +} + +func TestControl_StartStopLifecycle(t *testing.T) { + c, dev, conn := newReadyControl(t) + + require.NoError(t, c.Start()) + assert.Equal(t, StateStarted, c.State()) + require.ErrorIs(t, c.Start(), ErrAlreadyStarted) + + // Stop must unpark the reader blocked in the device and release everything + c.Stop() + assert.Equal(t, StateStopped, c.State()) + assert.True(t, dev.closed, "the tun device should have been closed") + assert.True(t, conn.closed, "the udp socket should have been closed") + require.ErrorIs(t, c.ctx.Err(), context.Canceled) + + // The reader drained off a closed device, that is not a fatal error + require.NoError(t, c.Wait()) + require.ErrorIs(t, c.Start(), ErrAlreadyStopped) +} + +func TestControl_RebindIsGatedByState(t *testing.T) { + c, _, conn := newReadyControl(t) + + // A rebind before Start reaches nothing, the interface is not up + c.RebindUDPServer() + assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op") + + require.NoError(t, c.Start()) + c.RebindUDPServer() + assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn") + + // A rebind racing a completed stop must not touch the closed conn + c.Stop() + require.NoError(t, c.Wait()) + c.RebindUDPServer() + assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op") +} diff --git a/interface.go b/interface.go index f96e431a..c44f38b3 100644 --- a/interface.go +++ b/interface.go @@ -215,6 +215,9 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) { ifce.connectionManager.intf = ifce + // Held until Close so waiting on the interface blocks until the resources are actually released + ifce.wg.Add(1) + return ifce, nil } @@ -258,17 +261,16 @@ func (f *Interface) activate() error { f.readers[i] = reader } - f.wg.Add(1) // for us to wait on Close() to return + // On error the caller owns the cleanup, Control.Start cancels the service context + // before releasing our resources so a waiter never observes a live context if err = f.inside.Activate(); err != nil { - f.wg.Done() - f.inside.Close() return err } return nil } -func (f *Interface) run() (func() error, error) { +func (f *Interface) run() { // Launch n queues to read packets from udp for i := 0; i < f.routines; i++ { f.wg.Go(func() { @@ -283,13 +285,14 @@ func (f *Interface) run() (func() error, error) { }) } - return func() error { - f.wg.Wait() - if e := f.fatalErr.Load(); e != nil { - return *e - } - return nil - }, nil +} + +func (f *Interface) wait() error { + f.wg.Wait() + if e := f.fatalErr.Load(); e != nil { + return *e + } + return nil } // onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one @@ -322,7 +325,10 @@ func (f *Interface) listenOut(i int) { f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get()) }) - if err != nil && !f.closed.Load() { + // An error after teardown began is shutdown noise, the closed flag covers resources + // Close releases itself and the cancelled ctx covers ones torn down by their owners + // reacting to it, like the user device pipes + if err != nil && !f.closed.Load() && f.ctx.Err() == nil { f.l.Error("Error while reading inbound packet, closing", "error", err) f.onFatal(err) } @@ -341,7 +347,8 @@ func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) { for { n, err := reader.Read(packet) if err != nil { - if !f.closed.Load() { + // Same shutdown noise handling as listenOut + if !f.closed.Load() && f.ctx.Err() == nil { f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i) f.onFatal(err) } @@ -542,9 +549,15 @@ func (f *Interface) GetCertState() *CertState { return f.pki.getCertState() } +// Close releases the interface's resources: the udp sockets and the tun device. +// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated, +// calls after the first return nil without doing anything. func (f *Interface) Close() error { + if !f.closed.CompareAndSwap(false, true) { + return nil + } + var errs []error - f.closed.Store(true) // Release the udp readers for i, u := range f.writers { @@ -560,6 +573,8 @@ func (f *Interface) Close() error { if closeErr != nil { errs = append(errs, closeErr) } + + // Release the construction token so waiters know the resources are gone f.wg.Done() return errors.Join(errs...) } diff --git a/main.go b/main.go index 7d7a0f72..d62d8dd0 100644 --- a/main.go +++ b/main.go @@ -130,6 +130,17 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev udpConns := make([]udp.Conn, routines) port := c.GetInt("listen.port", 0) + // Callers get no handle to these until the Control is returned, release them on any error. + defer func() { + if reterr != nil { + for _, u := range udpConns { + if u != nil { + _ = u.Close() + } + } + } + }() + if !configTest { rawListenHost := c.GetString("listen.host", "0.0.0.0") var listenHost netip.Addr diff --git a/overlay/tun_android.go b/overlay/tun_android.go index 9cbb64be..e4080b41 100644 --- a/overlay/tun_android.go +++ b/overlay/tun_android.go @@ -40,6 +40,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip err := t.reload(c, true) if err != nil { + _ = file.Close() return nil, err } diff --git a/overlay/tun_ios.go b/overlay/tun_ios.go index 6bfcbdfb..27bf558b 100644 --- a/overlay/tun_ios.go +++ b/overlay/tun_ios.go @@ -18,6 +18,7 @@ import ( "github.com/slackhq/nebula/config" "github.com/slackhq/nebula/routing" "github.com/slackhq/nebula/util" + "golang.org/x/sys/unix" ) type tun struct { @@ -33,6 +34,12 @@ func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error) } func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) { + if err := unix.SetNonblock(deviceFd, true); err != nil { + // We own the fd from the moment it is handed to us, same as the reload error path below + _ = unix.Close(deviceFd) + return nil, fmt.Errorf("failed to set the tun fd to non-blocking mode: %w", err) + } + file := os.NewFile(uintptr(deviceFd), "/dev/tun") t := &tun{ vpnNetworks: vpnNetworks, @@ -42,6 +49,7 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip err := t.reload(c, true) if err != nil { + _ = file.Close() return nil, err } diff --git a/service/service.go b/service/service.go index 899e851d..6610800d 100644 --- a/service/service.go +++ b/service/service.go @@ -43,12 +43,25 @@ type Service struct { } } -func New(control *nebula.Control) (*Service, error) { - wait, err := control.Start() +func New(control *nebula.Control) (_ *Service, reterr error) { + // Check this before Start so a failure doesn't leave a running nebula + device, ok := control.Device().(*overlay.UserDevice) + if !ok { + return nil, errors.New("must be using user device") + } + + err := control.Start() if err != nil { return nil, err } + // Anything that fails after a successful Start must tear nebula back down + defer func() { + if reterr != nil { + control.Stop() + } + }() + ctx := control.Context() eg, ctx := errgroup.WithContext(ctx) s := Service{ @@ -57,11 +70,6 @@ func New(control *nebula.Control) (*Service, error) { } s.mu.listeners = map[uint16]*tcpListener{} - device, ok := control.Device().(*overlay.UserDevice) - if !ok { - return nil, errors.New("must be using user device") - } - s.ipstack = stack.New(stack.Options{ NetworkProtocols: []stack.NetworkProtocolFactory{ipv4.NewProtocol, ipv6.NewProtocol}, TransportProtocols: []stack.TransportProtocolFactory{tcp.NewProtocol, udp.NewProtocol, icmp.NewProtocol4, icmp.NewProtocol6}, @@ -147,7 +155,7 @@ func New(control *nebula.Control) (*Service, error) { // Add the nebula wait function to the group so a fatal reader error // propagates out through errgroup.Wait(). eg.Go(func() error { - return wait() + return control.Wait() }) return &s, nil From 86733864fe7f13ad5664250af1c1726fe8275ab3 Mon Sep 17 00:00:00 2001 From: Jack Doan Date: Fri, 10 Jul 2026 12:02:17 -0500 Subject: [PATCH 21/25] don't make new relay state on a just-discarded tunnel (#1796) --- hostmap.go | 7 +++++-- relay_manager.go | 10 +++++++++- 2 files changed, 14 insertions(+), 3 deletions(-) diff --git a/hostmap.go b/hostmap.go index 2f2db101..b9acdd62 100644 --- a/hostmap.go +++ b/hostmap.go @@ -448,13 +448,15 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) { hm.unlockedMakePrimary(hostinfo) } -func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) { +// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses, +// false only when it is no longer in the hostmap at all. +func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool { // A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race // tunnel teardown, deciding to promote under the read lock and only taking the write lock // after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every // live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test. if hm.Indexes[hostinfo.localIndexId] != hostinfo { - return + return false } // Move hostinfo to the front (primary) of each of its address lists. The lists are @@ -469,6 +471,7 @@ func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) { list = append([]*HostInfo{hostinfo}, list...) hm.unlockedSetHostsForAddr(addr, list) } + return true } // unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index diff --git a/relay_manager.go b/relay_manager.go index 318a9f1a..1ae382a3 100644 --- a/relay_manager.go +++ b/relay_manager.go @@ -107,7 +107,10 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak if relayHostInfo.GetRemote().IsValid() { idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested) if err != nil { + // No local relay state was installed, so a CreateRelayRequest would hand the + // peer an index we could never resolve. Skip it. hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err) + continue } m := NebulaControl{ @@ -237,7 +240,12 @@ func AddRelay(l *slog.Logger, relayHostInfo *HostInfo, hm *HostMap, vpnIp netip. // Avoid standing up a relay that can't be used since only the primary hostinfo // will be pointed to by the relay logic //TODO: if there was an existing primary and it had relay state, should we merge? - hm.unlockedMakePrimary(relayHostInfo) + if !hm.unlockedMakePrimary(relayHostInfo) { + // The tunnel was torn down after the caller grabbed relayHostInfo. A relay standing + // on an unlinked hostinfo would never carry traffic, and its Relays entry could + // never be reclaimed since the delete-time cleanup has already run. + return 0, errors.New("relay hostinfo is no longer in the hostmap") + } hm.Relays[index] = relayHostInfo newRelay := Relay{ From 861d3aabd7186b2fa7108e92e346a53593bf24b7 Mon Sep 17 00:00:00 2001 From: Jack Doan Date: Mon, 13 Jul 2026 08:40:42 -0500 Subject: [PATCH 22/25] correct directionality of firewall.inbound_action and firewall.outbound_action (#1798) --- firewall.go | 16 ++++++++-------- inside.go | 4 ++-- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/firewall.go b/firewall.go index 84c505e7..f0fc79c9 100644 --- a/firewall.go +++ b/firewall.go @@ -44,8 +44,8 @@ type Firewall struct { InRules *FirewallTable OutRules *FirewallTable - InSendReject bool - OutSendReject bool + InboundSendReject bool + OutboundSendReject bool //TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better // https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt @@ -216,23 +216,23 @@ func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewal inboundAction := c.GetString("firewall.inbound_action", "drop") switch inboundAction { case "reject": - fw.InSendReject = true + fw.InboundSendReject = true case "drop": - fw.InSendReject = false + fw.InboundSendReject = false default: l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction) - fw.InSendReject = false + fw.InboundSendReject = false } outboundAction := c.GetString("firewall.outbound_action", "drop") switch outboundAction { case "reject": - fw.OutSendReject = true + fw.OutboundSendReject = true case "drop": - fw.OutSendReject = false + fw.OutboundSendReject = false default: l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction) - fw.OutSendReject = false + fw.OutboundSendReject = false } err := AddFirewallRulesFromConfig(l, false, c, fw) diff --git a/inside.go b/inside.go index fc079ba4..163a6034 100644 --- a/inside.go +++ b/inside.go @@ -87,7 +87,7 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet } func (f *Interface) rejectInside(packet []byte, out []byte, q int) { - if !f.firewall.InSendReject { + if !f.firewall.OutboundSendReject { return } @@ -103,7 +103,7 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) { } func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) { - if !f.firewall.OutSendReject { + if !f.firewall.InboundSendReject { return } From 6c3972f464df5a5fb17f7cd039b28fbfa55dd062 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Mon, 13 Jul 2026 11:49:59 -0500 Subject: [PATCH 23/25] code-sign: default the S3 key-prefix to the calling repo (#1799) --- .github/actions/code-sign/action.yml | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/.github/actions/code-sign/action.yml b/.github/actions/code-sign/action.yml index bfa1a9ec..f3956d95 100644 --- a/.github/actions/code-sign/action.yml +++ b/.github/actions/code-sign/action.yml @@ -25,9 +25,9 @@ inputs: required: false default: "code-signer" key-prefix: - description: "S3 key prefix the caller is authorized to write under" + description: "S3 key prefix to write under; defaults to code-signing// of the calling repo" required: false - default: "code-signing/slackhq/nebula" + default: "" runs: using: composite @@ -57,6 +57,9 @@ runs: KEY_PREFIX: ${{ inputs.key-prefix }} run: | set -eu + # Default the prefix to this repo so the S3 key attributes the sign correctly. + # nebula-nightly runs this same action but writes under its own repo's prefix. + KEY_PREFIX="${KEY_PREFIX:-code-signing/$GITHUB_REPOSITORY}" RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}" find "$SIGN_PATH" -name '*.exe' -print | while read -r path From e290a6892fee354b8de13ae1a23d28595e4cecda Mon Sep 17 00:00:00 2001 From: John Maguire Date: Fri, 17 Jul 2026 11:57:47 -0400 Subject: [PATCH 24/25] Fix relay re-establishment for handshake on Disestablised entry (#1805) handleOutsideRelayPacket filled ViaSender.remoteIdx with relay.RemoteIndex, an index from the relay peer's index space, but the rescue in sendHandshakeResponse looks that value up in relayForByIdx, which is keyed by local index. The lookup could never hit, so a terminal relay entry left Disestablished by a one-sided teardown stayed Disestablished even after a valid handshake arrived over it. The responder's first transmit then failed to find an Established relay, deleted its only relay entry, and every subsequent send was silently dropped until dead-tunnel detection forced a re-handshake. --- e2e/handshakes_test.go | 64 ++++++++++++++++++++++++++++++++++++++++++ handshake_manager.go | 2 +- hostmap.go | 1 - outside.go | 1 - 4 files changed, 65 insertions(+), 3 deletions(-) diff --git a/e2e/handshakes_test.go b/e2e/handshakes_test.go index eb5359cb..0c0bdf44 100644 --- a/e2e/handshakes_test.go +++ b/e2e/handshakes_test.go @@ -725,6 +725,70 @@ func TestReestablishRelays(t *testing.T) { } +func TestRelayHandshakeOverDisestablishedEntry(t *testing.T) { + t.Parallel() + // If them tears down the tunnel while me keeps Established relay state, me's next + // handshake flows through the relay with no fresh CreateRelayRequest and lands on + // them's Disestablished terminal relay entry. them must re-establish that entry, or + // its first transmit deletes its only relay and the tunnel is born transmit-dead: + // them can receive but every send is silently dropped. + ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) + myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}}) + relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}}) + theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}}) + + // Teach my how to get to the relay and that their can be reached via the relay + myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr) + myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()}) + relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr) + + // Build a router so we don't have to reason who gets which packet + r := router.NewR(t, myControl, relayControl, theirControl) + defer r.RenderFlow() + + // Start the servers + myControl.Start() + relayControl.Start() + theirControl.Start() + + t.Log("Trigger a handshake from me to them via the relay") + myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))) + + p := r.RouteForAllUntilTxTun(theirControl) + assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80) + oldIdx := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false).LocalIndex + + t.Log("Close the tunnel on them only, marking their relay entry Disestablished") + theirControl.CloseTunnel(myVpnIpNet[0].Addr(), true) + + t.Log("Re-handshake from me, riding the still-Established relay state") + myControl.ReHandshake(theirVpnIpNet[0].Addr()) + for { + h := myControl.GetHostInfoByVpnAddr(theirVpnIpNet[0].Addr(), false) + if h != nil && h.LocalIndex != oldIdx && h.RemoteIndex != 0 { + break + } + r.RouteForAllExitFunc(func(*udp.Packet, *nebula.Control) router.ExitType { + return router.RouteAndExit + }) + } + + hAtThem := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false) + require.NotNil(t, hAtThem, "them should have completed the relayed handshake") + require.Equal(t, []netip.Addr{relayVpnIpNet[0].Addr()}, hAtThem.CurrentRelaysToMe, "them should know a relay for the new tunnel") + + t.Log("Send from them to me; their only relay entry must survive the transmit") + theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))) + require.Never(t, func() bool { + h := theirControl.GetHostInfoByVpnAddr(myVpnIpNet[0].Addr(), false) + return h == nil || len(h.CurrentRelaysToMe) == 0 + }, time.Second, 10*time.Millisecond, "them deleted its only relay entry; the tunnel is permanently transmit-dead") + + p = r.RouteForAllUntilTxTun(myControl) + assertUdpPacket(t, []byte("Hi from them"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80) + r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl) +} + func TestStage1RaceRelays(t *testing.T) { t.Parallel() //NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay diff --git a/handshake_manager.go b/handshake_manager.go index 913918c2..6a2d0b4a 100644 --- a/handshake_manager.go +++ b/handshake_manager.go @@ -1077,7 +1077,7 @@ func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hos hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0]) // We received a valid handshake on this relay, so make sure the relay // state reflects that, in case it had been marked Disestablished. - via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established) + via.relayHI.relayState.UpdateRelayForByIdxState(via.relay.LocalIndex, Established) f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false) f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...) } diff --git a/hostmap.go b/hostmap.go index b9acdd62..45515fc3 100644 --- a/hostmap.go +++ b/hostmap.go @@ -287,7 +287,6 @@ type HostInfo struct { type ViaSender struct { UdpAddr netip.AddrPort relayHI *HostInfo // relayHI is the host info object of the relay - remoteIdx uint32 // remoteIdx is the index included in the header of the received packet relay *Relay // relay contains the rest of the relay information, including the PeerIP of the host trying to communicate with us. IsRelayed bool // IsRelayed is true if the packet was sent through a relay } diff --git a/outside.go b/outside.go index 4464acdf..8e89f807 100644 --- a/outside.go +++ b/outside.go @@ -214,7 +214,6 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, via = ViaSender{ UdpAddr: via.UdpAddr, relayHI: hostinfo, - remoteIdx: relay.RemoteIndex, relay: relay, IsRelayed: true, } From 147c202c27aa1cb934b19a384bbefee37a120ff5 Mon Sep 17 00:00:00 2001 From: Nate Brown Date: Fri, 17 Jul 2026 15:16:45 -0500 Subject: [PATCH 25/25] Swap back to a blocking udp socket, test `shutdown(2)` (#1806) Co-authored-by: Jack Doan --- cmd/nebula/close_on_timer_test.go | 96 +++++++++ udp/udp_linux.go | 335 +++++++++++++++--------------- udp/udp_linux_test.go | 179 ++++++++++++++++ 3 files changed, 443 insertions(+), 167 deletions(-) create mode 100644 cmd/nebula/close_on_timer_test.go create mode 100644 udp/udp_linux_test.go diff --git a/cmd/nebula/close_on_timer_test.go b/cmd/nebula/close_on_timer_test.go new file mode 100644 index 00000000..07138c15 --- /dev/null +++ b/cmd/nebula/close_on_timer_test.go @@ -0,0 +1,96 @@ +//go:build linux && !android && !e2e_testing + +package main + +import ( + "fmt" + "net/netip" + "os" + "path/filepath" + "runtime" + "testing" + "time" + + "github.com/slackhq/nebula" + "github.com/slackhq/nebula/cert" + cert_test "github.com/slackhq/nebula/cert_test" + "github.com/slackhq/nebula/config" + "github.com/slackhq/nebula/test" + "github.com/stretchr/testify/require" +) + +// TestControlStopClosesOnTimer reproduces the dnclient lifecycle: nebula runs as +// a library, and on a config update dnclient calls Stop() in-process to tear the +// old instance down before starting a new one. This boots a real nebula (real +// blocking UDP sockets, tun disabled), lets it run, then Stop()s it on a timer +// and asserts it actually closes. If the reader goroutines parked in recvmmsg +// don't wake on Close(), Wait() blocks forever and this fails with a goroutine +// dump instead of relying on a process signal to unstick them. +func TestControlStopClosesOnTimer(t *testing.T) { + l := test.NewLogger() + dir := t.TempDir() + + before := time.Now().Add(-time.Hour) + after := time.Now().Add(time.Hour) + ca, _, caKey, caPEM := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, before, after, nil, nil, nil) + networks := []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")} + _, _, keyPEM, certPEM := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, "close-on-timer", before, after, networks, nil, nil) + + caPath := filepath.Join(dir, "ca.pem") + certPath := filepath.Join(dir, "cert.pem") + keyPath := filepath.Join(dir, "key.pem") + require.NoError(t, os.WriteFile(caPath, caPEM, 0o600)) + require.NoError(t, os.WriteFile(certPath, certPEM, 0o600)) + require.NoError(t, os.WriteFile(keyPath, keyPEM, 0o600)) + + // tun disabled so no device/root is needed; routines: 2 so we exercise the + // multi-socket (SO_REUSEPORT) teardown, which is where dnclient runs. + configBody := fmt.Sprintf(` +pki: + ca: %s + cert: %s + key: %s +listen: + host: 127.0.0.1 + port: 0 +tun: + disabled: true +firewall: + outbound: + - port: any + proto: any + host: any + inbound: + - port: any + proto: any + host: any +routines: 2 +`, caPath, certPath, keyPath) + require.NoError(t, os.WriteFile(filepath.Join(dir, "config.yml"), []byte(configBody), 0o600)) + + c := config.NewC(l) + require.NoError(t, c.Load(dir)) + + ctrl, err := nebula.Main(c, false, "close-on-timer", l, nil) + require.NoError(t, err) + require.NoError(t, ctrl.Start()) + + // Run like a live nebula, then close on a timer, exactly as dnclient does. + <-time.NewTimer(5 * time.Second).C + + stopped := make(chan struct{}) + go func() { + ctrl.Stop() // closes the udp sockets (shutdown(2)) and the tun + ctrl.Wait() // blocks until every reader goroutine has returned + close(stopped) + }() + + select { + case <-stopped: + t.Log("nebula closed cleanly on timer") + case <-time.After(10 * time.Second): + buf := make([]byte, 1<<20) + n := runtime.Stack(buf, true) + t.Fatalf("nebula did NOT close within 10s of Stop(): a blocking reader never woke\n%s", buf[:n]) + } +} diff --git a/udp/udp_linux.go b/udp/udp_linux.go index 3e2d726a..3920342c 100644 --- a/udp/udp_linux.go +++ b/udp/udp_linux.go @@ -4,12 +4,13 @@ package udp import ( - "context" "encoding/binary" + "errors" "fmt" "log/slog" "net" "net/netip" + "sync/atomic" "syscall" "unsafe" @@ -19,58 +20,51 @@ import ( ) type StdConn struct { - udpConn *net.UDPConn - rawConn syscall.RawConn - isV4 bool - l *slog.Logger - batch int -} - -func setReusePort(network, address string, c syscall.RawConn) error { - var opErr error - err := c.Control(func(fd uintptr) { - opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1) - //CloseOnExec already set by the runtime - }) - if err != nil { - return err - } - return opErr + sysFd int + closed atomic.Bool + isV4 bool + l *slog.Logger + batch int } func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int) (Conn, error) { - listen := netip.AddrPortFrom(ip, uint16(port)) - lc := net.ListenConfig{} + af := unix.AF_INET6 + if ip.Is4() { + af = unix.AF_INET + } + syscall.ForkLock.RLock() + fd, err := unix.Socket(af, unix.SOCK_DGRAM, unix.IPPROTO_UDP) + if err == nil { + unix.CloseOnExec(fd) + } + syscall.ForkLock.RUnlock() + if err != nil { + return nil, fmt.Errorf("unable to open socket: %w", err) + } + if multi { - lc.Control = setReusePort - } - //this context is only used during the bind operation, you can't cancel it to kill the socket - pc, err := lc.ListenPacket(context.Background(), "udp", listen.String()) - if err != nil { - return nil, fmt.Errorf("unable to open socket: %s", err) - } - udpConn := pc.(*net.UDPConn) - rawConn, err := udpConn.SyscallConn() - if err != nil { - _ = udpConn.Close() - return nil, err - } - //gotta find out if we got an AF_INET6 socket or not: - out := &StdConn{ - udpConn: udpConn, - rawConn: rawConn, - l: l, - batch: batch, + if err = unix.SetsockoptInt(fd, unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil { + _ = unix.Close(fd) + return nil, fmt.Errorf("unable to set SO_REUSEPORT: %w", err) + } } - af, err := out.getSockOptInt(unix.SO_DOMAIN) - if err != nil { - _ = out.Close() - return nil, err + var sa unix.Sockaddr + if ip.Is4() { + sa4 := &unix.SockaddrInet4{Port: port} + sa4.Addr = ip.As4() + sa = sa4 + } else { + sa6 := &unix.SockaddrInet6{Port: port} + sa6.Addr = ip.As16() + sa = sa6 + } + if err = unix.Bind(fd, sa); err != nil { + _ = unix.Close(fd) + return nil, fmt.Errorf("unable to bind to socket: %w", err) } - out.isV4 = af == unix.AF_INET - return out, nil + return &StdConn{sysFd: fd, isV4: ip.Is4(), l: l, batch: batch}, nil } func (u *StdConn) SupportsMultipleReaders() bool { @@ -81,134 +75,111 @@ func (u *StdConn) Rebind() error { return nil } -func (u *StdConn) getSockOptInt(opt int) (int, error) { - if u.rawConn == nil { - return 0, fmt.Errorf("no UDP connection") - } - var out int - var opErr error - err := u.rawConn.Control(func(fd uintptr) { - out, opErr = unix.GetsockoptInt(int(fd), unix.SOL_SOCKET, opt) - }) - if err != nil { - return 0, err - } - return out, opErr -} - -func (u *StdConn) setSockOptInt(opt int, n int) error { - if u.rawConn == nil { - return fmt.Errorf("no UDP connection") - } - var opErr error - err := u.rawConn.Control(func(fd uintptr) { - opErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, opt, n) - }) - if err != nil { - return err - } - return opErr -} - func (u *StdConn) SetRecvBuffer(n int) error { - return u.setSockOptInt(unix.SO_RCVBUFFORCE, n) + return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUFFORCE, n) } func (u *StdConn) SetSendBuffer(n int) error { - return u.setSockOptInt(unix.SO_SNDBUFFORCE, n) + return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUFFORCE, n) } func (u *StdConn) SetSoMark(mark int) error { - return u.setSockOptInt(unix.SO_MARK, mark) + return unix.SetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK, mark) } func (u *StdConn) GetRecvBuffer() (int, error) { - return u.getSockOptInt(unix.SO_RCVBUF) + return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_RCVBUF) } func (u *StdConn) GetSendBuffer() (int, error) { - return u.getSockOptInt(unix.SO_SNDBUF) + return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_SNDBUF) } func (u *StdConn) GetSoMark() (int, error) { - return u.getSockOptInt(unix.SO_MARK) + return unix.GetsockoptInt(u.sysFd, unix.SOL_SOCKET, unix.SO_MARK) } func (u *StdConn) LocalAddr() (netip.AddrPort, error) { - a := u.udpConn.LocalAddr() - - switch v := a.(type) { - case *net.UDPAddr: - addr, ok := netip.AddrFromSlice(v.IP) - if !ok { - return netip.AddrPort{}, fmt.Errorf("LocalAddr returned invalid IP address: %s", v.IP) - } - return netip.AddrPortFrom(addr, uint16(v.Port)), nil - + sa, err := unix.Getsockname(u.sysFd) + if err != nil { + return netip.AddrPort{}, err + } + switch sa := sa.(type) { + case *unix.SockaddrInet4: + return netip.AddrPortFrom(netip.AddrFrom4(sa.Addr), uint16(sa.Port)), nil + case *unix.SockaddrInet6: + return netip.AddrPortFrom(netip.AddrFrom16(sa.Addr), uint16(sa.Port)), nil default: - return netip.AddrPort{}, fmt.Errorf("LocalAddr returned: %#v", a) + return netip.AddrPort{}, fmt.Errorf("unsupported sock type: %T", sa) } } -func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) { - var errno syscall.Errno - n, _, errno := unix.Syscall6( +// recvmmsg does one blocking recvmmsg (MSG_WAITFORONE), reading up to len(msgs) datagrams +func (u *StdConn) recvmmsg(msgs []rawMessage) (int, error) { + r, _, errno := unix.Syscall6( unix.SYS_RECVMMSG, - fd, + uintptr(u.sysFd), uintptr(unsafe.Pointer(&msgs[0])), uintptr(len(msgs)), unix.MSG_WAITFORONE, 0, 0, ) - if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK { - // No data available, block for I/O and try again. - return int(n), false, nil - } if errno != 0 { - return int(n), true, &net.OpError{Op: "recvmmsg", Err: errno} - } - return int(n), true, nil -} - -func (u *StdConn) listenOutSingle(r EncReader) error { - var err error - var n int - var from netip.AddrPort - buffer := make([]byte, MTU) - - for { - n, from, err = u.udpConn.ReadFromUDPAddrPort(buffer) - if err != nil { - return err + if u.closed.Load() { + return 0, net.ErrClosed } - from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port()) - r(from, buffer[:n]) + return 0, &net.OpError{Op: "recvmmsg", Err: errno} } + n := int(r) + if (n == 0 || msgs[0].Len == 0) && u.closed.Load() { + return 0, net.ErrClosed + } + return n, nil } -func (u *StdConn) listenOutBatch(r EncReader) error { +// recvmsg does one blocking recvmsg into msgs[0] +func (u *StdConn) recvmsg(msgs []rawMessage) (int, error) { + r, _, errno := unix.Syscall6( + unix.SYS_RECVMSG, + uintptr(u.sysFd), + uintptr(unsafe.Pointer(&msgs[0].Hdr)), + 0, + 0, + 0, + 0, + ) + if errno != 0 { + if u.closed.Load() { + return 0, net.ErrClosed + } + return 0, &net.OpError{Op: "recvmsg", Err: errno} + } + if r == 0 && u.closed.Load() { + return 0, net.ErrClosed + } + msgs[0].Len = uint32(r) + return 1, nil +} + +func (u *StdConn) ListenOut(r EncReader) error { var ip netip.Addr - var n int - var operr error - msgs, buffers, names := u.PrepareRawMessages(u.batch) - - //reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read - //defining it outside the loop so it gets re-used - reader := func(fd uintptr) (done bool) { - n, done, operr = recvmmsg(fd, msgs) - return done + read := u.recvmmsg + if u.batch == 1 { + read = u.recvmsg } for { - err := u.rawConn.Read(reader) + n, err := read(msgs) if err != nil { + if errors.Is(err, unix.EINTR) { + continue // interrupted by a signal, retry the read + } + // net.ErrClosed after Close() is teardown, absorbed by the caller's + // closed flag like the other platforms; anything else is a real error. return err } - if operr != nil { - return operr - } for i := 0; i < n; i++ { // Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic @@ -222,26 +193,68 @@ func (u *StdConn) listenOutBatch(r EncReader) error { } } -func (u *StdConn) ListenOut(r EncReader) error { - if u.batch == 1 { - return u.listenOutSingle(r) - } else { - return u.listenOutBatch(r) +func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error { + if u.isV4 { + return u.writeTo4(b, ip) + } + return u.writeTo6(b, ip) +} + +func (u *StdConn) writeTo6(b []byte, ip netip.AddrPort) error { + var rsa unix.RawSockaddrInet6 + rsa.Family = unix.AF_INET6 + rsa.Addr = ip.Addr().As16() + binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port()) + + for { + _, _, err := unix.Syscall6( + unix.SYS_SENDTO, + uintptr(u.sysFd), + uintptr(unsafe.Pointer(&b[0])), + uintptr(len(b)), + uintptr(0), + uintptr(unsafe.Pointer(&rsa)), + uintptr(unix.SizeofSockaddrInet6), + ) + if err != 0 { + return &net.OpError{Op: "sendto", Err: err} + } + return nil } } -func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error { - _, err := u.udpConn.WriteToUDPAddrPort(b, ip) - return err +func (u *StdConn) writeTo4(b []byte, ip netip.AddrPort) error { + if !ip.Addr().Is4() { + return ErrInvalidIPv6RemoteForSocket + } + + var rsa unix.RawSockaddrInet4 + rsa.Family = unix.AF_INET + rsa.Addr = ip.Addr().As4() + binary.BigEndian.PutUint16((*[2]byte)(unsafe.Pointer(&rsa.Port))[:], ip.Port()) + + for { + _, _, err := unix.Syscall6( + unix.SYS_SENDTO, + uintptr(u.sysFd), + uintptr(unsafe.Pointer(&b[0])), + uintptr(len(b)), + uintptr(0), + uintptr(unsafe.Pointer(&rsa)), + uintptr(unix.SizeofSockaddrInet4), + ) + if err != 0 { + return &net.OpError{Op: "sendto", Err: err} + } + return nil + } } func (u *StdConn) ReloadConfig(c *config.C) { b := c.GetInt("listen.read_buffer", 0) if b > 0 { - err := u.SetRecvBuffer(b) - if err == nil { - s, err := u.GetRecvBuffer() - if err == nil { + if err := u.SetRecvBuffer(b); err == nil { + if s, err := u.GetRecvBuffer(); err == nil { u.l.Info("listen.read_buffer was set", "size", s) } else { u.l.Warn("Failed to get listen.read_buffer", "error", err) @@ -253,10 +266,8 @@ func (u *StdConn) ReloadConfig(c *config.C) { b = c.GetInt("listen.write_buffer", 0) if b > 0 { - err := u.SetSendBuffer(b) - if err == nil { - s, err := u.GetSendBuffer() - if err == nil { + if err := u.SetSendBuffer(b); err == nil { + if s, err := u.GetSendBuffer(); err == nil { u.l.Info("listen.write_buffer was set", "size", s) } else { u.l.Warn("Failed to get listen.write_buffer", "error", err) @@ -269,10 +280,8 @@ func (u *StdConn) ReloadConfig(c *config.C) { b = c.GetInt("listen.so_mark", 0) s, err := u.GetSoMark() if b > 0 || (err == nil && s != 0) { - err := u.SetSoMark(b) - if err == nil { - s, err := u.GetSoMark() - if err == nil { + if err := u.SetSoMark(b); err == nil { + if s, err := u.GetSoMark(); err == nil { u.l.Info("listen.so_mark was set", "mark", s) } else { u.l.Warn("Failed to get listen.so_mark", "error", err) @@ -285,28 +294,20 @@ func (u *StdConn) ReloadConfig(c *config.C) { func (u *StdConn) getMemInfo(meminfo *[unix.SK_MEMINFO_VARS]uint32) error { var vallen uint32 = 4 * unix.SK_MEMINFO_VARS - - if u.rawConn == nil { - return fmt.Errorf("no UDP connection") - } - var opErr error - err := u.rawConn.Control(func(fd uintptr) { - _, _, syserr := unix.Syscall6(unix.SYS_GETSOCKOPT, fd, uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0) - if syserr != 0 { - opErr = syserr - } - }) - if err != nil { + _, _, err := unix.Syscall6(unix.SYS_GETSOCKOPT, uintptr(u.sysFd), uintptr(unix.SOL_SOCKET), uintptr(unix.SO_MEMINFO), uintptr(unsafe.Pointer(meminfo)), uintptr(unsafe.Pointer(&vallen)), 0) + if err != 0 { return err } - return opErr + return nil } func (u *StdConn) Close() error { - if u.udpConn != nil { - return u.udpConn.Close() - } - return nil + u.closed.Store(true) + // Wake the reader parked in recvmmsg/recvmsg. shutdown(2) on an unconnected socket + // returns ENOTCONN but still wakes it, so ignore the error. + // The reader then sees closed and stops touching the fd, making the Close below safe. + _ = unix.Shutdown(u.sysFd, unix.SHUT_RDWR) + return unix.Close(u.sysFd) } func NewUDPStatsEmitter(udpConns []Conn) func() { diff --git a/udp/udp_linux_test.go b/udp/udp_linux_test.go new file mode 100644 index 00000000..f9e7b3d8 --- /dev/null +++ b/udp/udp_linux_test.go @@ -0,0 +1,179 @@ +//go:build linux && !android && !e2e_testing + +package udp + +import ( + "errors" + "fmt" + "log/slog" + "net" + "net/netip" + "os" + "runtime" + "sync/atomic" + "testing" + "time" + + "golang.org/x/sys/unix" +) + +func testLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError})) +} + +// TestShutdownWakesAfterRx_Mechanism exercises the kernel quirk our teardown +// relies on: once a socket has received a packet, shutdown(2) wakes a blocked +// recvmmsg with n>=1/Len==0 (not n==0). recvmmsg must turn that into net.ErrClosed +// once Close set closed, so a parked reader exits instead of spinning. +func TestShutdownWakesAfterRx_Mechanism(t *testing.T) { + c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, 64) + if err != nil { + t.Fatalf("NewListener: %v", err) + } + sc := c.(*StdConn) + addr, err := sc.LocalAddr() + if err != nil { + t.Fatalf("LocalAddr: %v", err) + } + msgs, _, _ := sc.PrepareRawMessages(sc.batch) + + // Receive a real packet so the socket has carried data. + send, err := net.Dial("udp", addr.String()) + if err != nil { + t.Fatalf("dial: %v", err) + } + if _, err := send.Write([]byte("hello")); err != nil { + t.Fatalf("write: %v", err) + } + time.Sleep(50 * time.Millisecond) + n, err := sc.recvmmsg(msgs) + t.Logf("drain of real packet: n=%d err=%v msgs[0].Len=%d", n, err, msgs[0].Len) + _ = send.Close() + + // Block a reader on the now-empty queue, then tear down as Close() does. + // recvmmsg must return net.ErrClosed (not hang, not spin) even post-rx. + done := make(chan error, 1) + go func() { + _, err := sc.recvmmsg(msgs) + done <- err + }() + time.Sleep(150 * time.Millisecond) // let it park in recvmmsg + + sc.closed.Store(true) + if serr := unix.Shutdown(sc.sysFd, unix.SHUT_RDWR); serr != nil { + t.Logf("shutdown returned %v (expected ENOTCONN on unconnected UDP)", serr) + } + + select { + case err := <-done: + if !errors.Is(err, net.ErrClosed) { + t.Errorf("recvmmsg after post-rx shutdown returned %v, want net.ErrClosed", err) + } + case <-time.After(2 * time.Second): + t.Fatalf("HANG: recvmmsg did not return after shutdown following a received packet") + } + _ = unix.Close(sc.sysFd) +} + +// TestListenOutTeardown_TrafficPatterns reproduces the field report: a blocking +// reader must tear down cleanly on Close() regardless of what the socket has +// carried. The three cases the report called out: +// +// no traffic ever -> works (shutdown wakes recvmmsg with n==0) +// ping once, then idle -> historically HUNG: once the socket has received a +// packet, shutdown(2) wakes recvmmsg with n>=1/Len==0, +// which an n==0-only teardown check misses +// continuous traffic -> works (a real packet is always arriving) +// +// All three must return within the deadline; a hang dumps goroutines so the +// stuck reader is visible. +func TestListenOutTeardown_TrafficPatterns(t *testing.T) { + cases := []struct { + name string + traffic func(send net.Conn, stop <-chan struct{}) + }{ + {"no_traffic_ever", func(net.Conn, <-chan struct{}) {}}, + {"ping_once_then_idle", func(send net.Conn, _ <-chan struct{}) { + _, _ = send.Write([]byte("hello")) + }}, + {"continuous", func(send net.Conn, stop <-chan struct{}) { + for { + select { + case <-stop: + return + default: + _, _ = send.Write([]byte("hello")) + time.Sleep(2 * time.Millisecond) + } + } + }}, + } + + // batch 1 exercises the recvmsg path, batch 64 the recvmmsg path; both must + // tear down cleanly. + for _, batch := range []int{1, 64} { + for _, tc := range cases { + t.Run(fmt.Sprintf("batch%d/%s", batch, tc.name), func(t *testing.T) { + runTeardownCase(t, batch, tc.name, tc.traffic) + }) + } + } +} + +func runTeardownCase(t *testing.T, batch int, name string, traffic func(send net.Conn, stop <-chan struct{})) { + c, err := NewListener(testLogger(), netip.MustParseAddr("127.0.0.1"), 0, true, batch) + if err != nil { + t.Fatalf("NewListener: %v", err) + } + sc := c.(*StdConn) + addr, err := sc.LocalAddr() + if err != nil { + t.Fatalf("LocalAddr: %v", err) + } + + var received atomic.Int64 + loopDone := make(chan error, 1) + go func() { + loopDone <- sc.ListenOut(func(netip.AddrPort, []byte) { + received.Add(1) + }) + }() + + send, err := net.Dial("udp", addr.String()) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer send.Close() + + stop := make(chan struct{}) + trafficDone := make(chan struct{}) + go func() { + traffic(send, stop) + close(trafficDone) + }() + + // Let the pattern run and, for the idle case, the reader park again on an + // empty queue with the socket already having received a packet. + time.Sleep(500 * time.Millisecond) + + start := time.Now() + if err := sc.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + close(stop) + + select { + case err := <-loopDone: + // Clean teardown surfaces as net.ErrClosed (propagated like the other + // platforms); the caller absorbs it via its closed flag. + if err != nil && !errors.Is(err, net.ErrClosed) { + t.Fatalf("%s: ListenOut returned unexpected error on teardown: %v", name, err) + } + t.Logf("%s: closed in %v (received %d packets)", name, time.Since(start), received.Load()) + case <-time.After(3 * time.Second): + buf := make([]byte, 1<<20) + n := runtime.Stack(buf, true) + t.Fatalf("%s: HANG, ListenOut did not return within 3s of Close\n%s", name, buf[:n]) + } + <-trafficDone +}