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 diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 387963ea..ca2e45e0 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' }} @@ -163,14 +163,17 @@ 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 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 fbb24ecd..b15dff4e 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 391f3628..82d06385 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 40066574..447d4870 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: @@ -83,7 +83,7 @@ jobs: e2e-cmd: make e2evv steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -128,7 +128,7 @@ jobs: - {name: mobile, make-target: build-test-mobile} steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: 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") 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/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/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 ef58988b..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 @@ -305,7 +331,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 +376,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_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/control_test.go b/control_test.go index 5e381c46..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) { @@ -42,8 +45,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 +58,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 +78,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, @@ -119,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/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/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)) 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 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/e2e/handshakes_test.go b/e2e/handshakes_test.go index d0b9543c..0c0bdf44 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"))) @@ -721,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 @@ -819,18 +887,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 +992,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 +1097,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 +1191,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 +1291,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) @@ -1535,3 +1603,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/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") } diff --git a/firewall.go b/firewall.go index eb120fa6..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) @@ -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++ { diff --git a/go.mod b/go.mod index 804df03d..32aa5650 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 @@ -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..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= @@ -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= diff --git a/handshake_manager.go b/handshake_manager.go index 0d25305f..6a2d0b4a 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? @@ -1080,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 957894b6..45515fc3 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 } @@ -229,7 +238,7 @@ const ( ) type HostInfo struct { - remote netip.AddrPort + remote atomic.Pointer[netip.AddrPort] remotes *RemoteList promoteCounter atomic.Uint32 ConnectionState *ConnectionState @@ -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 @@ -282,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 } @@ -334,6 +338,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 +387,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 @@ -400,85 +447,66 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) { hm.unlockedMakePrimary(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 +// 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 false } - // If we are already primary then we won't bother re-linking - if oldHostinfo == 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) { + // 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 { - h := hm.Hosts[addr] - for h != nil { - if h == hostinfo { - hm.unlockedInnerDeleteHostInfo(h, addr) - } - h = h.next + if hm.Hosts[addr] == hostinfo { + // Already primary for this address, the list is already in the right order + continue } + list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo) + list = append([]*HostInfo{hostinfo}, list...) + hm.unlockedSetHostsForAddr(addr, list) } + return true } -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 +// 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 + } } } - hostinfo.next = nil - hostinfo.prev = nil + // 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{} + } + if len(hm.moreHosts) == 0 { + hm.moreHosts = map[netip.Addr][]*HostInfo{} + } // 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 @@ -502,7 +530,7 @@ func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Ad ) } - 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) @@ -511,6 +539,8 @@ func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Ad for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() { delete(hm.Relays, localRelayIdx) } + + return final } func (hm *HostMap) QueryIndex(index uint32) *HostInfo { @@ -554,19 +584,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") @@ -574,20 +615,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) } } } @@ -637,22 +672,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]) } } @@ -684,7 +724,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 +766,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 +789,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/hostmap_test.go b/hostmap_test.go index 2bd7bd43..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,243 @@ 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 +// 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 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. + 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)) + + // 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 +// 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 was pruned from both lists and the index map. + assert.Nil(t, hm.QueryIndex(oldest.localIndexId)) + + // 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) { diff --git a/inside.go b/inside.go index 27a6f758..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 } @@ -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/interface.go b/interface.go index 2aef678a..a89a6c12 100644 --- a/interface.go +++ b/interface.go @@ -216,6 +216,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 } @@ -262,17 +265,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() { @@ -287,13 +289,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 @@ -326,7 +329,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) } @@ -345,7 +351,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) } @@ -546,9 +553,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 { @@ -564,6 +577,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/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 diff --git a/lighthouse.go b/lighthouse.go index d23e84b8..3df74c39 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) @@ -1454,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 { @@ -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/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/outside.go b/outside.go index 7ebd4b5e..8e89f807 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 @@ -213,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, } @@ -276,7 +276,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) @@ -288,7 +289,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, ) } @@ -296,11 +297,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) } @@ -420,16 +421,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 { @@ -589,10 +588,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/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") +} 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_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_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) 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/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 { 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..1ae382a3 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,10 +104,13 @@ 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 { + // 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{ @@ -309,6 +317,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 +342,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 +424,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) @@ -508,7 +534,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 } 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) 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 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, 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 +}