diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index e4ca2933..9c5b1c3e 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -10,7 +10,7 @@ jobs: name: Build Linux/BSD All runs-on: ubuntu-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -36,7 +36,7 @@ jobs: id-token: write contents: read steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -76,7 +76,7 @@ jobs: HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }} runs-on: macos-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -134,7 +134,7 @@ jobs: # be overwritten - name: Checkout code if: ${{ env.HAS_DOCKER_CREDS == 'true' }} - uses: actions/checkout@v6 + uses: actions/checkout@v7 - name: Download artifacts if: ${{ env.HAS_DOCKER_CREDS == 'true' }} @@ -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 e0428e9c..8f71ead5 100644 --- a/.github/workflows/smoke-extra.yml +++ b/.github/workflows/smoke-extra.yml @@ -30,7 +30,7 @@ jobs: VAGRANT_DEFAULT_PROVIDER: libvirt steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -62,7 +62,7 @@ jobs: VAGRANT_DEFAULT_PROVIDER: virtualbox steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -88,7 +88,7 @@ jobs: runs-on: windows-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: diff --git a/.github/workflows/smoke.yml b/.github/workflows/smoke.yml index 4994a87e..ebac1cce 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: @@ -36,6 +36,14 @@ jobs: working-directory: ./.github/workflows/smoke run: ./smoke.sh + - name: setup docker image ipv6 + working-directory: ./.github/workflows/smoke + run: SMOKE_OVERLAY_IPV6=1 ./build.sh + + - name: run smoke ipv6 + working-directory: ./.github/workflows/smoke + run: SMOKE_OVERLAY_IPV6=1 ./smoke.sh + - name: setup relay docker image working-directory: ./.github/workflows/smoke run: ./build-relay.sh diff --git a/.github/workflows/smoke/build.sh b/.github/workflows/smoke/build.sh index b23516ee..fef76098 100755 --- a/.github/workflows/smoke/build.sh +++ b/.github/workflows/smoke/build.sh @@ -5,6 +5,19 @@ set -e -x rm -rf ./build mkdir ./build +if [ "$SMOKE_OVERLAY_IPV6" ] +then + LIGHTHOUSE_NIP="fd00:4242:0:0:0:ffff:c0a8:6401" + HOST2_NIP="fd00:4242:0:0:0:ffff:c0a8:6402" + HOST3_NIP="fd00:4242:0:0:0:ffff:c0a8:6403" + HOST4_NIP="fd00:4242:0:0:0:ffff:c0a8:6404" +else + LIGHTHOUSE_NIP="192.168.100.1" + HOST2_NIP="192.168.100.2" + HOST3_NIP="192.168.100.3" + HOST4_NIP="192.168.100.4" +fi + # Smoke containers run on a dedicated docker network whose subnet is allocated # at smoke time, not known at build time. Configs are written with TEST-NET-3 # placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh @@ -31,24 +44,24 @@ LIGHTHOUSE_IP="203.0.113.2" ../genconfig.sh >lighthouse1.yml HOST="host2" \ - LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \ + LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \ ../genconfig.sh >host2.yml HOST="host3" \ - LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \ + LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \ INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \ ../genconfig.sh >host3.yml HOST="host4" \ - LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \ + LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \ OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \ ../genconfig.sh >host4.yml ../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test" - ../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "192.168.100.1/24" - ../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "192.168.100.2/24" - ../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "192.168.100.3/24" - ../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24" + ../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "$LIGHTHOUSE_NIP/24" + ../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "$HOST2_NIP/24" + ../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "$HOST3_NIP/24" + ../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "$HOST4_NIP/24" ) docker build -t "nebula:${NAME:-smoke}" . diff --git a/.github/workflows/smoke/smoke.sh b/.github/workflows/smoke/smoke.sh index cad9dde7..f13ed380 100755 --- a/.github/workflows/smoke/smoke.sh +++ b/.github/workflows/smoke/smoke.sh @@ -47,6 +47,19 @@ HOST2_IP="$PREFIX.3" HOST3_IP="$PREFIX.4" HOST4_IP="$PREFIX.5" +if [ "$SMOKE_OVERLAY_IPV6" ] +then + LIGHTHOUSE_NIP="fd00:4242:0:0:0:ffff:c0a8:6401" + HOST2_NIP="fd00:4242:0:0:0:ffff:c0a8:6402" + HOST3_NIP="fd00:4242:0:0:0:ffff:c0a8:6403" + HOST4_NIP="fd00:4242:0:0:0:ffff:c0a8:6404" +else + LIGHTHOUSE_NIP="192.168.100.1" + HOST2_NIP="192.168.100.2" + HOST3_NIP="192.168.100.3" + HOST4_NIP="192.168.100.4" +fi + # Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones. # build/lighthouse1.yml has no IPs to rewrite so it's skipped. for f in build/host2.yml build/host3.yml build/host4.yml; do @@ -80,28 +93,28 @@ docker exec host3 tcpdump -i eth0 -q -w - -U 2>logs/host3.outside.log >logs/host docker exec host4 tcpdump -i tun0 -q -w - -U 2>logs/host4.inside.log >logs/host4.inside.pcap & docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap & -docker exec host2 ncat -nklv 0.0.0.0 2000 & -docker exec host3 ncat -nklv 0.0.0.0 2000 & -docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 0.0.0.0 4000 & -docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 0.0.0.0 3000 & -docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 & +docker exec host2 ncat -nklv 2000 & +docker exec host3 ncat -nklv 2000 & +docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 4000 & +docker exec host2 ncat -e '/usr/bin/echo host2' -nkluv 3000 & +docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 3000 & set +x echo echo " *** Testing ping from lighthouse1" echo set -x -docker exec lighthouse1 ping -c1 192.168.100.2 -docker exec lighthouse1 ping -c1 192.168.100.3 +docker exec lighthouse1 ping -c1 $HOST2_NIP +docker exec lighthouse1 ping -c1 $HOST3_NIP set +x echo echo " *** Testing ping from host2" echo set -x -docker exec host2 ping -c1 192.168.100.1 +docker exec host2 ping -c1 $LIGHTHOUSE_NIP # Should fail because not allowed by host3 inbound firewall -! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1 +! docker exec host2 ping -c1 $HOST3_NIP -w5 || exit 1 set +x echo @@ -109,34 +122,34 @@ echo " *** Testing ncat from host2" echo set -x # Should fail because not allowed by host3 inbound firewall -! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1 -! docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1 +! docker exec host2 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1 +! docker exec host2 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1 set +x echo echo " *** Testing ping from host3" echo set -x -docker exec host3 ping -c1 192.168.100.1 -docker exec host3 ping -c1 192.168.100.2 +docker exec host3 ping -c1 $LIGHTHOUSE_NIP +docker exec host3 ping -c1 $HOST2_NIP set +x echo echo " *** Testing ncat from host3" echo set -x -docker exec host3 ncat -nzv -w5 192.168.100.2 2000 -docker exec host3 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 +docker exec host3 ncat -nzv -w5 $HOST2_NIP 2000 +docker exec host3 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2 set +x echo echo " *** Testing ping from host4" echo set -x -docker exec host4 ping -c1 192.168.100.1 +docker exec host4 ping -c1 $LIGHTHOUSE_NIP # Should fail because not allowed by host4 outbound firewall -! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1 -! docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1 +! docker exec host4 ping -c1 $HOST2_NIP -w5 || exit 1 +! docker exec host4 ping -c1 $HOST3_NIP -w5 || exit 1 set +x echo @@ -144,10 +157,10 @@ echo " *** Testing ncat from host4" echo set -x # Should fail because not allowed by host4 outbound firewall -! docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1 -! docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || exit 1 -! docker exec host4 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 || exit 1 -! docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1 +! docker exec host4 ncat -nzv -w5 $HOST2_NIP 2000 || exit 1 +! docker exec host4 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1 +! docker exec host4 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2 || exit 1 +! docker exec host4 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1 set +x echo @@ -159,7 +172,7 @@ set -x # cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000, # conntrack must let host4's listener reply on that flow. If it doesn't, # the echo back from host4 never reaches host2. -docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4 +docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv $HOST4_NIP 4000" | grep -q helloagainfromhost4 docker exec host4 sh -c 'kill 1' docker exec host3 sh -c 'kill 1' diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 2abb3740..269f0edb 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -18,7 +18,7 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -78,7 +78,7 @@ jobs: e2e-cmd: make e2evv steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: @@ -123,7 +123,7 @@ jobs: - {name: mobile, make-target: build-test-mobile} steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@v7 - uses: actions/setup-go@v6 with: diff --git a/Makefile b/Makefile index 7fefac7e..e2216d48 100644 --- a/Makefile +++ b/Makefile @@ -272,6 +272,9 @@ smoke-multiport-docker: bin-docker cd .github/workflows/smoke/ && NAME="smoke-multiport" MULTIPORT_TX=true MULTIPORT_RX=true MULTIPORT_HANDSHAKE=true ./build.sh cd .github/workflows/smoke/ && NAME="smoke-multiport" ./smoke.sh +smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1 +smoke-docker-ipv6: smoke-docker + smoke-docker-race: BUILD_ARGS = -race smoke-docker-race: CGO_ENABLED = 1 smoke-docker-race: smoke-docker 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/connection_manager.go b/connection_manager.go index ee6d1eaf..88f31321 100644 --- a/connection_manager.go +++ b/connection_manager.go @@ -136,14 +136,6 @@ func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time) return in, out } -// AddTrafficWatch must be called for every new HostInfo. -// We will continue to monitor the HostInfo until the tunnel is dropped. -func (cm *connectionManager) AddTrafficWatch(h *HostInfo) { - if h.out.Swap(true) == false { - cm.trafficTimer.Add(h.localIndexId, cm.checkInterval) - } -} - func (cm *connectionManager) Start(ctx context.Context) { clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration) defer clockSource.Stop() @@ -306,8 +298,8 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo } else { cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu)) cm.l.Info("send CreateRelayRequest", - "relayFrom", req.RelayFromAddr, - "relayTo", req.RelayToAddr, + "relayFrom", relayFrom, + "relayTo", relayTo, "initiatorRelayIndex", req.InitiatorRelayIndex, "responderRelayIndex", req.ResponderRelayIndex, "vpnAddrs", newhostinfo.vpnAddrs, diff --git a/control.go b/control.go index ef58988b..053feab5 100644 --- a/control.go +++ b/control.go @@ -305,7 +305,7 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) { c.l.Debug("Sending close tunnel message", "vpnAddrs", h.vpnAddrs, - "udpAddr", h.remote, + "udpAddr", h.GetRemote(), ) closed++ } @@ -350,7 +350,7 @@ func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo { RemoteAddrs: h.remotes.CopyAddrs(preferredRanges), CurrentRelaysToMe: h.relayState.CopyRelayIps(), CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(), - CurrentRemote: h.remote, + CurrentRemote: h.GetRemote(), } for i, a := range h.vpnAddrs { diff --git a/control_test.go b/control_test.go index 5e381c46..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/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..d580eb21 100644 --- a/e2e/handshakes_test.go +++ b/e2e/handshakes_test.go @@ -1535,3 +1535,78 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) { myControl.Stop() theirControl.Stop() } + +func TestMultiVpnAddrDeletePrimaryKeepsSecondAddr(t *testing.T) { + t.Parallel() + // Regression for the hostmap multi-vpnAddr delete bug. A dual-stack (v4+v6) V2-cert peer that + // handshakes twice at once ends up with two hostinfos linked in the shared next/prev chain, with the + // primary owning both addresses. Deleting that primary (e.g. connection manager dropping it, a + // CloseTunnel, a collision) must promote the surviving sibling for EVERY address. The pre-fix code + // unlinked the chain once per address, so it promoted the sibling for the first address and orphaned + // the second: the peer stayed reachable at its v4 addr but not its v6 addr despite a live tunnel. + + ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) + myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fd00::1/64", nil) + theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "10.128.0.2/24,fd00::2/64", nil) + + // This bug only exists for peers carrying more than one vpn address + require.Len(t, theirVpnIpNet, 2) + theirV4 := theirVpnIpNet[0].Addr() + theirV6 := theirVpnIpNet[1].Addr() + + // Put their info in our lighthouse and vice versa + myControl.InjectLightHouseAddr(theirV4, theirUdpAddr) + theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr) + + // Build a router so we don't have to reason who gets which packet + r := router.NewR(t, myControl, theirControl) + defer r.RenderFlow() + + myControl.Start() + theirControl.Start() + + // Race a handshake so both of us build a hostinfo for the other, leaving my hostmap with a single + // host (them) backed by two linked hostinfos, just like TestStage1Race. + myControl.InjectTunPacket(BuildTunUDPPacket(theirV4, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))) + theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirV4, 80, []byte("Hi from them"))) + + myHsForThem := myControl.GetFromUDP(true) + theirHsForMe := theirControl.GetFromUDP(true) + + r.InjectUDPPacket(theirControl, myControl, theirHsForMe) + r.InjectUDPPacket(myControl, theirControl, myHsForThem) + + r.RouteForAllUntilTxTun(theirControl) + r.RouteForAllUntilTxTun(myControl) + + r.RenderHostmaps("Racing hostmaps", myControl, theirControl) + + // Two hostinfos for them means the shared next/prev chain has a sibling to promote. The Hosts map has + // one entry per vpn address (two, for dual stack), so the index count is what tells us there are two + // hostinfos. + require.Len(t, myControl.ListHostmapIndexes(false), 2) + + // The primary owns both of their addresses + primaryV4 := myControl.GetHostInfoByVpnAddr(theirV4, false) + primaryV6 := myControl.GetHostInfoByVpnAddr(theirV6, false) + require.NotNil(t, primaryV4) + require.NotNil(t, primaryV6) + require.Equal(t, primaryV4.LocalIndex, primaryV6.LocalIndex, "both addrs should point at the same primary") + + // Delete the primary tunnel. localOnly so we don't perturb their side, we only care about my hostmap. + require.True(t, myControl.CloseTunnel(theirV4, true)) + + // The surviving sibling must still serve BOTH addresses. + survivorV4 := myControl.GetHostInfoByVpnAddr(theirV4, false) + survivorV6 := myControl.GetHostInfoByVpnAddr(theirV6, false) + require.NotNil(t, survivorV4, "v4 addr should still resolve to the surviving tunnel") + // Pre-fix this is nil: the second address was orphaned when the primary was deleted. + require.NotNil(t, survivorV6, "v6 addr was orphaned after deleting the primary (multi-vpnAddr delete bug)") + assert.Equal(t, survivorV4.LocalIndex, survivorV6.LocalIndex, "both addrs should promote to the same survivor") + assert.NotEqual(t, primaryV4.LocalIndex, survivorV4.LocalIndex, "a different hostinfo should now be primary") + + r.RenderHostmaps("Final hostmaps", myControl, theirControl) + + myControl.Stop() + theirControl.Stop() +} diff --git a/e2e/tunnels_test.go b/e2e/tunnels_test.go index 697f25af..18c69a3f 100644 --- a/e2e/tunnels_test.go +++ b/e2e/tunnels_test.go @@ -15,6 +15,7 @@ import ( "github.com/slackhq/nebula/header" "github.com/slackhq/nebula/udp" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "gopkg.in/yaml.v3" ) @@ -373,6 +374,100 @@ func TestCrossStackRelaysWork(t *testing.T) { //relayControl.Stop() } +// TestRelayReplayProtection asserts that a relay (forwarding-type) node rejects +// replayed relay frames. A captured relay frame, re-injected with the same +// message counter, must be dropped by the replay window rather than re-forwarded +// to the relay target. Before the fix, handleOutsideRelayPacket authenticated the +// frame but never advanced the replay window, so every replay was re-forwarded. +func TestRelayReplayProtection(t *testing.T) { + t.Parallel() + ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) + myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version2, ca, caKey, "me ", "10.128.0.1/24,fc00::1/64", m{"relay": m{"use_relays": true}}) + relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}}) + theirUdp := netip.MustParseAddrPort("10.0.0.2:4242") + theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdp(cert.Version2, ca, caKey, "them ", "fc00::2/64", theirUdp, m{"relay": m{"use_relays": true}}) + + myVpnV6 := myVpnIpNet[1] + relayVpnV4 := relayVpnIpNet[0] + relayVpnV6 := relayVpnIpNet[1] + theirVpnV6 := theirVpnIpNet[0] + + // Teach me how to reach the relay and that them is reachable via the relay + myControl.InjectLightHouseAddr(relayVpnV4.Addr(), relayUdpAddr) + myControl.InjectLightHouseAddr(relayVpnV6.Addr(), relayUdpAddr) + myControl.InjectRelays(theirVpnV6.Addr(), []netip.Addr{relayVpnV6.Addr()}) + relayControl.InjectLightHouseAddr(theirVpnV6.Addr(), theirUdpAddr) + + r := router.NewR(t, myControl, relayControl, theirControl) + defer r.RenderFlow() + + myControl.Start() + relayControl.Start() + theirControl.Start() + + // Establish the relayed tunnel in both directions so all handshakes complete. + t.Log("Establish the relayed tunnel") + myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me"))) + p := r.RouteForAllUntilTxTun(theirControl) + assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80) + theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them"))) + p = r.RouteForAllUntilTxTun(myControl) + assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80) + + // Drain anything still queued on me's UDP tx so the next packet we pull is the + // relay frame we are about to generate. + for myControl.GetFromUDP(false) != nil { + } + + // Capture a single legitimate relay frame that me transmits toward the relay. + t.Log("Capture a relay frame from me -> relay") + myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("replay me"))) + relayFrame := myControl.GetFromUDP(true) + require.Equal(t, relayUdpAddr, relayFrame.To, "captured frame should be addressed to the relay") + var fh header.H + require.NoError(t, fh.Parse(relayFrame.Data)) + require.Equal(t, header.Message, fh.Type) + require.Equal(t, header.MessageRelay, fh.Subtype) + + // drainForwards counts relay frames the relay forwards toward them within the + // settle window. We match on destination + (Message, MessageRelay) so the + // relay's own direct traffic to them can't be miscounted. + drainForwards := func(settle time.Duration) int { + ch := relayControl.GetUDPTxChan() + count := 0 + for { + select { + case pkt := <-ch: + var ph header.H + if pkt.To == theirUdpAddr && ph.Parse(pkt.Data) == nil && + ph.Type == header.Message && ph.Subtype == header.MessageRelay { + count++ + } + pkt.Release() + case <-time.After(settle): + return count + } + } + } + + // First delivery of the captured frame: the relay should forward it once. + t.Log("Deliver the captured frame once; relay forwards it to them") + relayControl.InjectUDPPacket(relayFrame) + require.Equal(t, 1, drainForwards(200*time.Millisecond), "relay should forward the first, legitimate copy") + + // Replay the exact same frame several times. A correct replay window rejects + // these duplicates so the relay forwards none of them. + t.Log("Replay the captured frame; relay must drop the duplicates") + const replays = 3 + for i := 0; i < replays; i++ { + relayControl.InjectUDPPacket(relayFrame) + } + forwarded := drainForwards(200 * time.Millisecond) + assert.Equal(t, 0, forwarded, "relay re-forwarded %d/%d replayed relay frames; replay protection is ineffective on relay tunnels", forwarded, replays) + + r.RenderHostmaps("Final hostmaps", myControl, relayControl, theirControl) +} + func TestCloseTunnelAuthenticated(t *testing.T) { t.Parallel() ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) diff --git a/examples/config.yml b/examples/config.yml index 679f0e9b..309dcbfd 100644 --- a/examples/config.yml +++ b/examples/config.yml @@ -438,7 +438,7 @@ firewall: # `drop` (default): silently drop the packet. # `reject`: send a reject reply. # - For TCP, this will be a RST "Connection Reset" packet. - # - For other protocols, this will be an ICMP port unreachable packet. + # - For other protocols, this will be an ICMP "Destination unreachable: Communication administratively prohibited" packet. outbound_action: drop inbound_action: drop diff --git a/firewall.go b/firewall.go index eb120fa6..84c505e7 100644 --- a/firewall.go +++ b/firewall.go @@ -423,11 +423,6 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table") // Drop returns an error if the packet should be dropped, explaining why. It // returns nil if the packet should not be dropped. func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error { - // Check if we spoke to this tuple, if we did then allow this packet - if f.inConns(fp, h, caPool, localCache) { - return nil - } - // Make sure remote address matches nebula certificate, and determine how to treat it if h.networks == nil { // Simple case: Certificate has one address and no unsafe networks @@ -461,6 +456,11 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool * return ErrInvalidLocalIP } + // Check if we spoke to this tuple, if we did then allow this packet + if f.inConns(fp, h, caPool, localCache) { + return nil + } + table := f.OutRules if incoming { table = f.InRules diff --git a/firewall_test.go b/firewall_test.go index 9373f1fd..499f3cc7 100644 --- a/firewall_test.go +++ b/firewall_test.go @@ -916,6 +916,159 @@ func TestFirewall_DropIPSpoofing(t *testing.T) { assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP) } +func TestFirewall_ConntrackSourceSpoofingAcrossPeers(t *testing.T) { + l := test.NewLoggerWithOutput(&bytes.Buffer{}) + + myVpnNetworksTable := new(bart.Lite) + myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24")) + + owner := &dummyCert{ + name: "owner", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")}, + } + + victim := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "victim", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")}, + }, + } + victimHI := HostInfo{ + ConnectionState: &ConnectionState{peerCert: victim}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")}, + } + victimHI.buildNetworks(myVpnNetworksTable, victim.Certificate) + + attacker := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "attacker", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.3/24")}, + }, + } + attackerHI := HostInfo{ + ConnectionState: &ConnectionState{peerCert: attacker}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.3")}, + } + attackerHI.buildNetworks(myVpnNetworksTable, attacker.Certificate) + + fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner) + // Allow any inbound traffic that passes the cert / source-IP checks. + require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + cp := cert.NewCAPool() + + flow := firewall.Packet{ + LocalAddr: netip.MustParseAddr("192.0.2.1"), + RemoteAddr: netip.MustParseAddr("192.0.2.2"), + LocalPort: 443, + RemotePort: 55000, + Protocol: firewall.ProtoUDP, + } + + require.NoError(t, fw.Drop(flow, true, &victimHI, cp, nil), + "victim's own traffic from its own overlay IP must be allowed") + + unseen := flow + unseen.RemotePort = 55001 + assert.Equal(t, ErrInvalidRemoteIP, fw.Drop(unseen, true, &attackerHI, cp, nil), + "sanity: attacker forging victim's source IP must be rejected when no conntrack entry exists") + + got := fw.Drop(flow, true, &attackerHI, cp, nil) + t.Logf("attacker replaying victim's 4-tuple: Drop returned %v (nil == packet ALLOWED == spoof succeeded)", got) + assert.Equal(t, ErrInvalidRemoteIP, got, + "SECURITY: attacker spoofed victim's overlay source IP (192.0.2.2) by reusing an existing conntrack 4-tuple; Drop returned %v instead of rejecting", got) +} + +// BenchmarkFirewallDropConntrackHit measures Drop on an already-established flow +// (a conntrack hit). This is the fast path that the source-IP<->cert binding +// reordering adds work to, so it quantifies the cost of moving the address checks +// ahead of the conntrack lookup. Cases: +// - simple: peer cert has one address, no unsafe networks (h.networks == nil), +// so the remote-address check is a single netip.Addr compare. +// - complex: peer cert has unsafe networks (h.networks populated), so the +// remote-address check is a BART lookup. +// - noCache/localCache: whether a per-batch ConntrackCache is supplied, which in +// the original code let the fast path skip straight past the address checks. +func BenchmarkFirewallDropConntrackHit(b *testing.B) { + l := test.NewLoggerWithOutput(&bytes.Buffer{}) + + myVpnNetworksTable := new(bart.Lite) + myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24")) + + owner := &dummyCert{ + name: "owner", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")}, + } + + simpleCert := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "simple", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")}, + }, + } + simpleHost := &HostInfo{ + ConnectionState: &ConnectionState{peerCert: simpleCert}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")}, + } + simpleHost.buildNetworks(myVpnNetworksTable, simpleCert.Certificate) + + complexCert := &cert.CachedCertificate{ + Certificate: &dummyCert{ + name: "complex", + networks: []netip.Prefix{netip.MustParsePrefix("192.0.2.2/24")}, + unsafeNetworks: []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}, + }, + } + complexHost := &HostInfo{ + ConnectionState: &ConnectionState{peerCert: complexCert}, + vpnAddrs: []netip.Addr{netip.MustParseAddr("192.0.2.2")}, + } + complexHost.buildNetworks(myVpnNetworksTable, complexCert.Certificate) + + cp := cert.NewCAPool() + + flow := firewall.Packet{ + LocalAddr: netip.MustParseAddr("192.0.2.1"), + RemoteAddr: netip.MustParseAddr("192.0.2.2"), + LocalPort: 443, + RemotePort: 55000, + Protocol: firewall.ProtoUDP, + } + + cases := []struct { + name string + host *HostInfo + useCache bool + }{ + {"simple/noCache", simpleHost, false}, + {"simple/localCache", simpleHost, true}, + {"complex/noCache", complexHost, false}, + {"complex/localCache", complexHost, true}, + } + + for _, tc := range cases { + b.Run(tc.name, func(b *testing.B) { + fw := NewFirewall(l, time.Second, time.Minute, time.Hour, owner) + require.NoError(b, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"any"}, "", "", "", "", "")) + + // Establish the conntrack entry so every benchmarked Drop is a hit. + require.NoError(b, fw.Drop(flow, true, tc.host, cp, nil)) + + var cache firewall.ConntrackCache + if tc.useCache { + cache = firewall.ConntrackCache{} + } + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := fw.Drop(flow, true, tc.host, cp, cache); err != nil { + b.Fatal(err) + } + } + }) + } +} + func BenchmarkLookup(b *testing.B) { ml := func(m map[string]struct{}, a [][]string) { for n := 0; n < b.N; n++ { diff --git a/go.mod b/go.mod index bd1c0c57..b30a97a1 100644 --- a/go.mod +++ b/go.mod @@ -9,10 +9,10 @@ require ( github.com/armon/go-radix v1.0.0 github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432 github.com/flynn/noise v1.1.0 - github.com/gaissmai/bart v0.27.1 + 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 @@ -24,15 +24,15 @@ require ( github.com/vishvananda/netlink v1.3.1 go.uber.org/goleak v1.3.0 go.yaml.in/yaml/v3 v3.0.4 - golang.org/x/crypto v0.51.0 + golang.org/x/crypto v0.53.0 golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 - golang.org/x/net v0.54.0 - golang.org/x/sync v0.20.0 - golang.org/x/sys v0.44.0 - golang.org/x/term v0.43.0 + golang.org/x/net v0.56.0 + golang.org/x/sync v0.21.0 + golang.org/x/sys v0.46.0 + 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 8ab36d34..11e72276 100644 --- a/go.sum +++ b/go.sum @@ -26,8 +26,8 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg= github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag= -github.com/gaissmai/bart v0.27.1 h1:FysPzqETMJa8q9rNkLW5peT1hq25nLOz8ksHbSVoiAk= -github.com/gaissmai/bart v0.27.1/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c= +github.com/gaissmai/bart v0.28.0 h1:89yZLo8NmyqD0RYgJ3QO9HhqqGGw+oWhf90cZm69Lko= +github.com/gaissmai/bart v0.28.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c= github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as= github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as= github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY= @@ -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= @@ -162,16 +162,16 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= -golang.org/x/crypto v0.51.0 h1:IBPXwPfKxY7cWQZ38ZCIRPI50YLeevDLlLnyC5wRGTI= -golang.org/x/crypto v0.51.0/go.mod h1:8AdwkbraGNABw2kOX6YFPs3WM22XqI4EXEd8g+x7Oc8= +golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= +golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY= golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc= golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY= 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= @@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA= golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= -golang.org/x/net v0.54.0 h1:2zJIZAxAHV/OHCDTCOHAYehQzLfSXuf/5SoL/Dv6w/w= -golang.org/x/net v0.54.0/go.mod h1:Sj4oj8jK6XmHpBZU/zWHw3BV3abl4Kvi+Ut7cQcY+cQ= +golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= +golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= @@ -191,8 +191,8 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= -golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= +golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= @@ -208,11 +208,11 @@ golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.44.0 h1:ildZl3J4uzeKP07r2F++Op7E9B29JRUy+a27EibtBTQ= -golang.org/x/sys v0.44.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= +golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= -golang.org/x/term v0.43.0 h1:S4RLU2sB31O/NCl+zFN9Aru9A/Cq2aqKpTZJ6B+DwT4= -golang.org/x/term v0.43.0/go.mod h1:lrhlHNdQJHO+1qVYiHfFKVuVioJIheAc3fBSMFYEIsk= +golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc= +golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= @@ -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/errors.go b/handshake/errors.go index bb8a5893..3bdcc947 100644 --- a/handshake/errors.go +++ b/handshake/errors.go @@ -13,6 +13,7 @@ var ( ErrUnknownSubtype = errors.New("unknown handshake subtype") ErrMissingContent = errors.New("expected handshake content but message was empty") ErrUnexpectedContent = errors.New("received unexpected handshake content") + ErrInvalidRemoteIndex = errors.New("peer sent an invalid index in handshake payload") ErrIndexAllocation = errors.New("failed to allocate local index") ErrNoCredential = errors.New("no handshake credential available for cert version") ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key") diff --git a/handshake/machine.go b/handshake/machine.go index 42dd48f4..393f528b 100644 --- a/handshake/machine.go +++ b/handshake/machine.go @@ -323,8 +323,9 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error { // Process payload if flags.expectsPayload { + var remoteIndex uint32 if m.result.Initiator { - m.result.RemoteIndex = payload.ResponderIndex + remoteIndex = payload.ResponderIndex if payload.ResponderMultiPort != nil { m.result.MultiportRx = payload.ResponderMultiPort.RxSupported m.result.MultiportTx = payload.ResponderMultiPort.TxSupported @@ -333,7 +334,7 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error { } } } else { - m.result.RemoteIndex = payload.InitiatorIndex + remoteIndex = payload.InitiatorIndex if payload.InitiatorMultiPort != nil { m.result.MultiportRx = payload.InitiatorMultiPort.RxSupported m.result.MultiportTx = payload.InitiatorMultiPort.TxSupported @@ -342,6 +343,13 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error { } } } + // The payload presence check above can be satisfied by Time alone, so a payload + // could still carry a zero index here. We need to reject it. + if remoteIndex == 0 { + m.failed = true + return ErrInvalidRemoteIndex + } + m.result.RemoteIndex = remoteIndex m.result.HandshakeTime = payload.Time m.payloadSet = true } diff --git a/handshake/machine_test.go b/handshake/machine_test.go index 5818cc3b..3a0e7a85 100644 --- a/handshake/machine_test.go +++ b/handshake/machine_test.go @@ -230,6 +230,24 @@ func TestMachineProcessPayload(t *testing.T) { require.ErrorIs(t, err, ErrUnexpectedContent) assert.True(t, m.Failed()) }) + + t.Run("zero initiator index on responder is fatal", func(t *testing.T) { + m := newTestMachine(t, cs, v, false, 100) + bytes := MarshalPayload(nil, Payload{InitiatorIndex: 0, Time: 1}) + err := m.processPayload(bytes, msgFlags{expectsPayload: true}) + require.ErrorIs(t, err, ErrInvalidRemoteIndex) + assert.True(t, m.Failed()) + assert.Zero(t, m.result.RemoteIndex) + }) + + t.Run("zero responder index on initiator is fatal", func(t *testing.T) { + m := newTestMachine(t, cs, v, true, 100) + bytes := MarshalPayload(nil, Payload{InitiatorIndex: 100, ResponderIndex: 0, Time: 1}) + err := m.processPayload(bytes, msgFlags{expectsPayload: true}) + require.ErrorIs(t, err, ErrInvalidRemoteIndex) + assert.True(t, m.Failed()) + assert.Zero(t, m.result.RemoteIndex) + }) } // TestMachineRequireComplete checks the fail-on-incomplete-handshake path diff --git a/handshake_manager.go b/handshake_manager.go index bc5efda3..7253a04e 100644 --- a/handshake_manager.go +++ b/handshake_manager.go @@ -837,7 +837,6 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head } hm.sendHandshakeResponse(via, response, hostinfo, false) - f.connectionManager.AddTrafficWatch(hostinfo) hostinfo.remotes.RefreshFromHandshake(vpnAddrs) // Don't wait for UpdateWorker @@ -1014,7 +1013,6 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate) hm.Complete(hostinfo, f) - f.connectionManager.AddTrafficWatch(hostinfo) if len(hh.packetStore) > 0 { if f.l.Enabled(context.Background(), slog.LevelDebug) { @@ -1152,7 +1150,7 @@ func (hm *HandshakeManager) handleCheckAndCompleteError(err error, existing, hos case ErrAlreadySeen: if hostinfo.multiportRx { // The other host is sending to us with multiport, so only grab the IP - via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), hostinfo.remote.Port()) + via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), hostinfo.GetRemote().Port()) } if existing.SetRemoteIfPreferred(f.hostMap, via) { f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu)) diff --git a/hostmap.go b/hostmap.go index 4ffe319d..8a44902f 100644 --- a/hostmap.go +++ b/hostmap.go @@ -138,9 +138,9 @@ func (rs *RelayState) InsertRelayTo(ip netip.Addr) { } func (rs *RelayState) CopyRelayIps() []netip.Addr { - ret := make([]netip.Addr, len(rs.relays)) rs.RLock() defer rs.RUnlock() + ret := make([]netip.Addr, len(rs.relays)) copy(ret, rs.relays) return ret } @@ -229,7 +229,7 @@ const ( ) type HostInfo struct { - remote netip.AddrPort + remote atomic.Pointer[netip.AddrPort] remotes *RemoteList promoteCounter atomic.Uint32 ConnectionState *ConnectionState @@ -444,43 +444,29 @@ func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) { } func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) { + isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil + for _, addr := range hostinfo.vpnAddrs { - h := hm.Hosts[addr] - for h != nil { - if h == hostinfo { - hm.unlockedInnerDeleteHostInfo(h, addr) - } - h = h.next + if hm.Hosts[addr] != hostinfo { + continue + } + if hostinfo.next != nil { + // Promote the next hostinfo in the shared chain to primary for this address + hm.Hosts[addr] = hostinfo.next + } else { + delete(hm.Hosts, addr) } } -} + if len(hm.Hosts) == 0 { + hm.Hosts = map[netip.Addr]*HostInfo{} + } -func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Addr) { - primary, ok := hm.Hosts[addr] - isLastHostinfo := hostinfo.next == nil && hostinfo.prev == nil - if ok && primary == hostinfo { - // The vpn addr pointer points to the same hostinfo as the local index id, we can remove it - delete(hm.Hosts, addr) - if len(hm.Hosts) == 0 { - hm.Hosts = map[netip.Addr]*HostInfo{} - } - - if hostinfo.next != nil { - // We had more than 1 hostinfo at this vpn addr, promote the next in the list to primary - hm.Hosts[addr] = hostinfo.next - // It is primary, there is no previous hostinfo now - hostinfo.next.prev = nil - } - - } else { - // Relink if we were in the middle of multiple hostinfos for this vpn addr - if hostinfo.prev != nil { - hostinfo.prev.next = hostinfo.next - } - - if hostinfo.next != nil { - hostinfo.next.prev = hostinfo.prev - } + // Splice this hostinfo out of the shared chain exactly once + if hostinfo.prev != nil { + hostinfo.prev.next = hostinfo.next + } + if hostinfo.next != nil { + hostinfo.next.prev = hostinfo.prev } hostinfo.next = nil @@ -629,6 +615,11 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) { hm.Indexes[hostinfo.localIndexId] = hostinfo hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo + hostinfo.out.Store(true) + if f.connectionManager != nil { // f.connectionManager is only nil in some unit tests + f.connectionManager.trafficTimer.Add(hostinfo.localIndexId, f.connectionManager.checkInterval) + } + if hm.l.Enabled(context.Background(), slog.LevelDebug) { hm.l.Debug("Hostmap vpnIp added", "hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts), @@ -685,7 +676,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() { @@ -727,11 +718,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) } } @@ -743,7 +741,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..156444a3 100644 --- a/hostmap_test.go +++ b/hostmap_test.go @@ -194,6 +194,107 @@ func TestHostMap_DeleteHostInfo(t *testing.T) { assert.Nil(t, prim) } +// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one +// vpnAddr and shares its next/prev chain with a live sibling. Deleting the head must not corrupt the +// sibling: every address the sibling owns has to keep pointing at it. The pre-fix code unlinked the shared +// chain once per vpnAddr, so on the first address it nil'd next/prev, and on the second address the node +// looked already-detached: it dropped the map entry instead of promoting the sibling (and tripped the +// isLastHostinfo relay teardown). See unlockedDeleteHostInfo. +func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + + f := &Interface{} + + a := netip.MustParseAddr("0.0.0.1") + b := netip.MustParseAddr("0.0.0.2") + + // Two tunnels for the same peer, each reachable at both a and b. + other := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 1} + head := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2} + + hm.unlockedAddHostInfo(other, f) + hm.unlockedAddHostInfo(head, f) + + // head is primary for both addresses, other is next in the shared chain + assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(a).localIndexId) + assert.Equal(t, head.localIndexId, hm.QueryVpnAddr(b).localIndexId) + assert.Equal(t, other.localIndexId, head.next.localIndexId) + assert.Equal(t, head.localIndexId, other.prev.localIndexId) + + // Delete the head. other is still live, so it must become primary for BOTH addresses. + hm.DeleteHostInfo(head) + + // Pre-fix: QueryVpnAddr(b) came back nil here because the second address was deleted rather than + // promoted, leaving other unreachable at b. + require.NotNil(t, hm.QueryVpnAddr(a)) + require.NotNil(t, hm.QueryVpnAddr(b)) + assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(a).localIndexId) + assert.Equal(t, other.localIndexId, hm.QueryVpnAddr(b).localIndexId) + + // other is now the only hostinfo in the chain + assert.Nil(t, other.prev) + assert.Nil(t, other.next) + + // head is fully detached + assert.Nil(t, head.prev) + assert.Nil(t, head.next) + assert.Nil(t, hm.QueryIndex(head.localIndexId)) +} + +// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune +// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long) +// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is +// primary for none of the addresses, and both address chains must stay consistent afterwards. +func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) { + l := test.NewLogger() + hm := newHostMap(l) + + f := &Interface{} + + a := netip.MustParseAddr("0.0.0.1") + b := netip.MustParseAddr("0.0.0.2") + + // Add one more than the cap, newest last so it becomes head. Every hostinfo owns both a and b. + hostinfos := make([]*HostInfo, 0, MaxHostInfosPerVpnIp+1) + for i := 0; i <= MaxHostInfosPerVpnIp; i++ { + hostinfos = append(hostinfos, &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: uint32(i + 1)}) + } + // Add oldest first (highest index in our slice) so the very first one added is the overflow victim. + for i := len(hostinfos) - 1; i >= 0; i-- { + hm.unlockedAddHostInfo(hostinfos[i], f) + } + + oldest := hostinfos[len(hostinfos)-1] + + // The oldest hostinfo should have been pruned and fully detached + assert.Nil(t, oldest.next) + assert.Nil(t, oldest.prev) + assert.Nil(t, hm.QueryIndex(oldest.localIndexId)) + + // Both addresses resolve to the same head, and that head is one of the survivors (not the pruned one) + primA := hm.QueryVpnAddr(a) + primB := hm.QueryVpnAddr(b) + require.NotNil(t, primA) + require.NotNil(t, primB) + assert.Equal(t, primA.localIndexId, primB.localIndexId) + assert.NotEqual(t, oldest.localIndexId, primA.localIndexId) + + // Walk the shared chain: exactly MaxHostInfosPerVpnIp survivors, no cycles, oldest absent + seen := map[uint32]struct{}{} + for h := primA; h != nil; h = h.next { + _, dup := seen[h.localIndexId] + require.False(t, dup, "cycle detected in hostinfo chain") + seen[h.localIndexId] = struct{}{} + if h.next != nil { + assert.Equal(t, h.localIndexId, h.next.prev.localIndexId, "prev pointer must mirror next") + } + } + assert.Len(t, seen, MaxHostInfosPerVpnIp) + _, prunedStillPresent := seen[oldest.localIndexId] + assert.False(t, prunedStillPresent) +} + func TestHostMap_reload(t *testing.T) { l := test.NewLogger() c := config.NewC(test.NewLogger()) diff --git a/inside.go b/inside.go index 75a50a71..1737d3d7 100644 --- a/inside.go +++ b/inside.go @@ -334,7 +334,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) } @@ -362,7 +362,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType } } - useRelay := !remote.IsValid() && !hostinfo.remote.IsValid() + useRelay := !remote.IsValid() && !hostinfo.GetRemote().IsValid() fullOut := out if useRelay { @@ -427,13 +427,13 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType "udpAddr", remote, ) } - } else if hostinfo.remote.IsValid() { + } else if hr := hostinfo.GetRemote(); hr.IsValid() { if multiport { rawOut = rawOut[:len(out)+udp.RawOverhead] port := udpPortGetter.UDPSendPort(f.multiPort.TxPorts) - err = f.udpRaw.WriteTo(rawOut, port, hostinfo.remote) + err = f.udpRaw.WriteTo(rawOut, port, hr) } else { - err = f.writers[q].WriteTo(out, hostinfo.remote) + err = f.writers[q].WriteTo(out, hr) } if err != nil { hostinfo.logger(f.l).Error("Failed to write outgoing packet", diff --git a/iputil/packet.go b/iputil/packet.go index b18e5244..c0c1921e 100644 --- a/iputil/packet.go +++ b/iputil/packet.go @@ -4,26 +4,54 @@ import ( "encoding/binary" "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" ) const ( - // Need 96 bytes for the largest reject packet: + // MaxIPv4RejectPacketSize is the largest IPv4 reject packet: // - 20 byte ipv4 header // - 8 byte icmpv4 header // - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header) - MaxRejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8 + maxIPv4RejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8 + + // MaxRejectPacketSize is sized for the largest possible reject packet (IPv6): + // - 40 byte ipv6 header + // - 8 byte icmpv6 header + // - up to 1000 byte body (original packet, possibly truncated. We want to stay + // under the MTU with Nebula overhead included) + maxIPv6RejectPacketSize = ipv6.HeaderLen + 8 + 1000 + + MaxRejectPacketSize = maxIPv6RejectPacketSize ) func CreateRejectPacket(packet []byte, out []byte) []byte { - if len(packet) < ipv4.HeaderLen || int(packet[0]>>4) != ipv4.Version { + if len(packet) < 1 { return nil } - switch packet[9] { - case 6: // tcp - return ipv4CreateRejectTCPPacket(packet, out) + version := int(packet[0] >> 4) + switch version { + case ipv4.Version: + if len(packet) < ipv4.HeaderLen { + return nil + } + // Do not send reject packets for non-first fragments + if packet[6]&0x1f != 0 || packet[7] != 0 { + return nil + } + switch packet[9] { + case 6: // tcp + return ipv4CreateRejectTCPPacket(packet, out) + default: + return ipv4CreateRejectICMPPacket(packet, out) + } + case ipv6.Version: + if len(packet) < ipv6.HeaderLen { + return nil + } + return ipv6CreateRejectPacket(packet, out) default: - return ipv4CreateRejectICMPPacket(packet, out) + return nil } } @@ -35,12 +63,17 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte { return nil } - // ICMP reply includes original header and first 8 bytes of the packet - packetLen := len(packet) - if packetLen > ihl+8 { - packetLen = ihl + 8 + // Do not generate ICMP errors in response to ICMP error packets + if packet[9] == 1 && len(packet) > ihl { + icmpType := packet[ihl] + if icmpType == 3 || icmpType == 4 || icmpType == 5 || icmpType == 11 || icmpType == 12 { + return nil + } } + // ICMP reply includes original header and first 8 bytes of the packet + packetLen := min(len(packet), ihl+8) + outLen := ipv4.HeaderLen + 8 + packetLen if outLen > cap(out) { return nil @@ -71,14 +104,14 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte { // ICMP Destination Unreachable icmpOut := out[ipv4.HeaderLen:] - icmpOut[0] = 3 // type (Destination unreachable) - icmpOut[1] = 3 // code (Port unreachable error) - icmpOut[2] = 0 // checksum - icmpOut[3] = 0 // . - icmpOut[4] = 0 // unused - icmpOut[5] = 0 // . - icmpOut[6] = 0 // . - icmpOut[7] = 0 // . + icmpOut[0] = 3 // type (Destination unreachable) + icmpOut[1] = 13 // code (Communication administratively prohibited) + icmpOut[2] = 0 // checksum + icmpOut[3] = 0 // . + icmpOut[4] = 0 // unused + icmpOut[5] = 0 // . + icmpOut[6] = 0 // . + icmpOut[7] = 0 // . // Copy original IP header and first 8 bytes as body copy(icmpOut[8:], packet[:packetLen]) @@ -165,7 +198,193 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte { return out } +func ipv6CreateRejectPacket(packet []byte, out []byte) []byte { + proto, offset, isFragment := ipv6FindUpperProtocol(packet) + if isFragment { + return nil + } + switch proto { + case 6: // tcp + return ipv6CreateRejectTCPPacket(packet, out, offset) + default: + return ipv6CreateRejectICMPPacket(packet, out, proto, offset) + } +} + +func ipv6CreateRejectICMPPacket(packet []byte, out []byte, proto uint8, offset int) []byte { + // Do not generate ICMPv6 errors in response to ICMPv6 error packets + if proto == 58 && len(packet) > offset { + icmpType := packet[offset] + if icmpType >= 1 && icmpType <= 4 { + return nil + } + } + + // Include as much of the original packet as possible, up to 1000 bytes, + // so the response fits comfortably within any tunnel MTU. + packetLen := min(len(packet), 1000) + + outLen := ipv6.HeaderLen + 8 + packetLen + if outLen > cap(out) { + return nil + } + + out = out[:outLen] + + // IPv6 header + ipHdr := out[0:ipv6.HeaderLen] + ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits) + ipHdr[1] = 0 // traffic class (low bits), flow label (high bits) + ipHdr[2] = 0 // flow label + ipHdr[3] = 0 // flow label + + payloadLen := uint16(outLen - ipv6.HeaderLen) + binary.BigEndian.PutUint16(ipHdr[4:], payloadLen) // payload length + ipHdr[6] = 58 // next header (ICMPv6) + ipHdr[7] = 64 // hop limit + + // Swap dest / src IPs (each 16 bytes, src at 8, dst at 24) + copy(ipHdr[8:24], packet[24:40]) + copy(ipHdr[24:40], packet[8:24]) + + // ICMPv6 Destination Unreachable + icmpOut := out[ipv6.HeaderLen:] + icmpOut[0] = 1 // type (Destination Unreachable) + icmpOut[1] = 1 // code (Communication with destination administratively prohibited) + icmpOut[2] = 0 // checksum + icmpOut[3] = 0 // . + icmpOut[4] = 0 // unused + icmpOut[5] = 0 // . + icmpOut[6] = 0 // . + icmpOut[7] = 0 // . + + copy(icmpOut[8:], packet[:packetLen]) + + // ICMPv6 checksum uses a pseudo-header + csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 58, uint32(payloadLen)) + binary.BigEndian.PutUint16(icmpOut[2:], tcpipChecksum(icmpOut, csum)) + + return out +} + +func ipv6CreateRejectTCPPacket(packet []byte, out []byte, offset int) []byte { + const tcpLen = 20 + + if len(packet) < offset+tcpLen { + return nil + } + + outLen := ipv6.HeaderLen + tcpLen + if outLen > cap(out) { + return nil + } + + out = out[:outLen] + + // IPv6 header + ipHdr := out[0:ipv6.HeaderLen] + ipHdr[0] = ipv6.Version << 4 // version, traffic class (high bits) + ipHdr[1] = 0 // traffic class (low bits), flow label (high bits) + ipHdr[2] = 0 // flow label + ipHdr[3] = 0 // flow label + + binary.BigEndian.PutUint16(ipHdr[4:], tcpLen) // payload length + ipHdr[6] = 6 // next header (TCP) + ipHdr[7] = 64 // hop limit + + // Swap dest / src IPs + copy(ipHdr[8:24], packet[24:40]) + copy(ipHdr[24:40], packet[8:24]) + + // TCP RST + tcpIn := packet[offset:] + var ackSeq, seq uint32 + outFlags := byte(0b00000100) // RST + + inAck := tcpIn[13]&0b00010000 != 0 + if inAck { + seq = binary.BigEndian.Uint32(tcpIn[8:]) + } else { + inSyn := uint32((tcpIn[13] & 0b00000010) >> 1) + inFin := uint32(tcpIn[13] & 0b00000001) + ackSeq = binary.BigEndian.Uint32(tcpIn[4:]) + inSyn + inFin + uint32(len(tcpIn)) - uint32(tcpIn[12]>>4)<<2 + outFlags |= 0b00010000 // ACK + } + + tcpOut := out[ipv6.HeaderLen:] + // Swap dest / src ports + copy(tcpOut[0:2], tcpIn[2:4]) + copy(tcpOut[2:4], tcpIn[0:2]) + binary.BigEndian.PutUint32(tcpOut[4:], seq) + binary.BigEndian.PutUint32(tcpOut[8:], ackSeq) + tcpOut[12] = (tcpLen >> 2) << 4 // data offset, reserved, NS + tcpOut[13] = outFlags // CWR, ECE, URG, ACK, PSH, RST, SYN, FIN + tcpOut[14] = 0 // window size + tcpOut[15] = 0 // . + tcpOut[16] = 0 // checksum + tcpOut[17] = 0 // . + tcpOut[18] = 0 // URG Pointer + tcpOut[19] = 0 // . + + // Calculate checksum with IPv6 pseudo-header + csum := ipv6PseudoheaderChecksum(ipHdr[8:24], ipHdr[24:40], 6, tcpLen) + binary.BigEndian.PutUint16(tcpOut[16:], tcpipChecksum(tcpOut, csum)) + + return out +} + +func ipv6FindUpperProtocol(packet []byte) (nextHeader uint8, offset int, isFragment bool) { + nextHeader = packet[6] + offset = ipv6.HeaderLen + + for { + switch nextHeader { + case 0, 43, 60: // Hop-by-Hop, Routing, Destination + if len(packet) < offset+2 { + return nextHeader, offset, isFragment + } + nextHeader = packet[offset] + offset += (int(packet[offset+1]) + 1) << 3 + + case 44: // Fragment + if len(packet) < offset+8 { + return nextHeader, offset, isFragment + } + if packet[offset+2] != 0 || packet[offset+3]&0xf8 != 0 { + isFragment = true + } + nextHeader = packet[offset] + offset += 8 + + case 51: // AH + if len(packet) < offset+2 { + return nextHeader, offset, isFragment + } + nextHeader = packet[offset] + offset += (int(packet[offset+1]) + 2) << 2 + + default: + return nextHeader, offset, isFragment + } + } +} + func CreateICMPEchoResponse(packet, out []byte) []byte { + if len(packet) < 1 { + return nil + } + + switch packet[0] >> 4 { + case 4: + return createICMPv4EchoResponse(packet, out) + case 6: + return createICMPv6EchoResponse(packet, out) + default: + return nil + } +} + +func createICMPv4EchoResponse(packet, out []byte) []byte { // Return early if this is not a simple ICMP Echo Request //TODO: make constants out of these if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) { @@ -199,6 +418,43 @@ func CreateICMPEchoResponse(packet, out []byte) []byte { return out } +func createICMPv6EchoResponse(packet, out []byte) []byte { + // IPv6 header (40 bytes) + ICMPv6 header (8 bytes minimum) + if len(packet) < ipv6.HeaderLen+8 || len(packet) > 9001 { + return nil + } + + // Next Header must be ICMPv6 (58) + if packet[6] != 58 { + return nil + } + + // ICMPv6 type must be Echo Request (128) + if packet[ipv6.HeaderLen] != 128 { + return nil + } + + out = out[:len(packet)] + copy(out, packet) + + // Swap src/dst addresses (bytes 8-23 and 24-39) + copy(out[8:24], packet[24:40]) + copy(out[24:40], packet[8:24]) + + // Change ICMPv6 type to Echo Reply (129) + icmp := out[ipv6.HeaderLen:] + icmp[0] = 129 + icmp[2] = 0 + icmp[3] = 0 + + // ICMPv6 checksum uses a pseudo-header with src, dst, length, and next header + payloadLen := uint32(len(icmp)) + csum := ipv6PseudoheaderChecksum(out[8:24], out[24:40], 58, payloadLen) + binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum)) + + return out +} + // calculates the TCP/IP checksum defined in rfc1071. The passed-in // csum is any initial checksum data that's already been computed. // @@ -236,3 +492,18 @@ func ipv4PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint3 csum += length >> 16 return csum } + +// based on: +// - https://github.com/google/gopacket/blob/v1.1.19/layers/tcpip.go#L37-L48 +func ipv6PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint32) { + for i := 0; i < 16; i += 2 { + csum += uint32(src[i]) << 8 + csum += uint32(src[i+1]) + csum += uint32(dst[i]) << 8 + csum += uint32(dst[i+1]) + } + csum += proto + csum += length & 0xffff + csum += length >> 16 + return csum +} diff --git a/iputil/packet_test.go b/iputil/packet_test.go index e1d0d95d..6d567d51 100644 --- a/iputil/packet_test.go +++ b/iputil/packet_test.go @@ -1,11 +1,13 @@ package iputil import ( + "encoding/binary" "net" "testing" "github.com/stretchr/testify/assert" "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" ) func Test_CreateRejectPacket(t *testing.T) { @@ -43,7 +45,7 @@ func Test_CreateRejectPacket(t *testing.T) { } b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...) - expectedLen = MaxRejectPacketSize + expectedLen = maxIPv4RejectPacketSize out = make([]byte, MaxRejectPacketSize) rejectPacket = CreateRejectPacket(b, out) assert.NotNil(t, rejectPacket) @@ -71,3 +73,404 @@ func Test_CreateRejectPacket(t *testing.T) { assert.NotNil(t, rejectPacket) assert.Len(t, rejectPacket, expectedLen) } + +func Test_CreateRejectPacket_NoFragment(t *testing.T) { + out := make([]byte, MaxRejectPacketSize) + + // IPv4: non-zero fragment offset should not generate reject packet + h := ipv4.Header{ + Len: 20, + Src: net.IPv4(10, 0, 0, 1), + Dst: net.IPv4(10, 0, 0, 2), + Protocol: 17, // UDP + } + b, err := h.Marshal() + if err != nil { + t.Fatalf("h.Marshal: %v", err) + } + b = append(b, make([]byte, 8)...) + // Set fragment offset to non-zero (byte 6-7, offset in 8-byte units) + b[6] = 0x00 + b[7] = 0x01 + assert.Nil(t, CreateRejectPacket(b, out)) + + // MF flag with zero offset (first fragment) should still generate reject + b[6] = 0x20 // MF flag set + b[7] = 0x00 + assert.NotNil(t, CreateRejectPacket(b, out)) + + // Non-fragment should still generate reject packet + b[6] = 0x00 + b[7] = 0x00 + assert.NotNil(t, CreateRejectPacket(b, out)) + + // DF flag only (not a fragment) should still generate reject packet + b[6] = 0x40 + b[7] = 0x00 + assert.NotNil(t, CreateRejectPacket(b, out)) +} + +func Test_CreateRejectPacketIPv6_NoFragment(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + out := make([]byte, MaxRejectPacketSize) + + // IPv6 with Fragment header and non-zero offset should not generate reject + fragHeader := []byte{ + 17, // next header: UDP + 0, // reserved + 0, 9, // fragment offset=1 (shifted left 3), M=1 + 0, 0, 0, 1, // identification + } + udpPayload := make([]byte, 8) + payload := append(fragHeader, udpPayload...) + packet := makeIPv6Packet(src, dst, 44, payload) // next header 44 = Fragment + assert.Nil(t, CreateRejectPacket(packet, out)) + + // Fragment header with zero offset (first fragment) should still generate reject + fragHeader[2] = 0 + fragHeader[3] = 1 // offset=0, M=1 + payload = append(fragHeader, udpPayload...) + packet = makeIPv6Packet(src, dst, 44, payload) + assert.NotNil(t, CreateRejectPacket(packet, out)) +} + +func Test_CreateRejectPacket_NoICMPError(t *testing.T) { + out := make([]byte, MaxRejectPacketSize) + + // ICMP error types should not generate reject packets + icmpErrorTypes := []byte{3, 4, 5, 11, 12} + for _, icmpType := range icmpErrorTypes { + h := ipv4.Header{ + Len: 20, + Src: net.IPv4(10, 0, 0, 1), + Dst: net.IPv4(10, 0, 0, 2), + Protocol: 1, // ICMP + } + + b, err := h.Marshal() + if err != nil { + t.Fatalf("h.Marshal: %v", err) + } + b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0) + + rejectPacket := CreateRejectPacket(b, out) + assert.Nil(t, rejectPacket, "ICMP type %d should not generate a reject packet", icmpType) + } + + // ICMP non-error types should still generate reject packets + icmpNonErrorTypes := []byte{0, 8, 13, 14} + for _, icmpType := range icmpNonErrorTypes { + h := ipv4.Header{ + Len: 20, + Src: net.IPv4(10, 0, 0, 1), + Dst: net.IPv4(10, 0, 0, 2), + Protocol: 1, // ICMP + } + + b, err := h.Marshal() + if err != nil { + t.Fatalf("h.Marshal: %v", err) + } + b = append(b, icmpType, 0, 0, 0, 0, 0, 0, 0) + + rejectPacket := CreateRejectPacket(b, out) + assert.NotNil(t, rejectPacket, "ICMP type %d should generate a reject packet", icmpType) + } +} + +func makeIPv6Packet(src, dst net.IP, nextHeader uint8, payload []byte) []byte { + b := make([]byte, ipv6.HeaderLen+len(payload)) + b[0] = ipv6.Version << 4 + binary.BigEndian.PutUint16(b[4:], uint16(len(payload))) + b[6] = nextHeader + b[7] = 64 + copy(b[8:24], src.To16()) + copy(b[24:40], dst.To16()) + copy(b[ipv6.HeaderLen:], payload) + return b +} + +func Test_CreateRejectPacketIPv6_ICMP(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + + // Small UDP packet: entire original included in body + udpPayload := make([]byte, 20) + udpPayload[0] = 0x00 // src port high + udpPayload[1] = 0x50 // src port low (80) + udpPayload[2] = 0x01 // dst port high + udpPayload[3] = 0xBB // dst port low (443) + packet := makeIPv6Packet(src, dst, 17, udpPayload) + + out := make([]byte, MaxRejectPacketSize) + rejectPacket := CreateRejectPacket(packet, out) + assert.NotNil(t, rejectPacket) + + // Small packet fits entirely: 40 (ipv6 hdr) + 8 (icmpv6 hdr) + 60 (original) + expectedLen := ipv6.HeaderLen + 8 + len(packet) + assert.Len(t, rejectPacket, expectedLen) + + // Verify version + assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0) + // Verify next header is ICMPv6 (58) + assert.Equal(t, byte(58), rejectPacket[6]) + // Verify src/dst are swapped + assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24])) + assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40])) + // Verify ICMPv6 type=1 (Dest Unreachable), code=1 (Administratively prohibited) + assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen]) + assert.Equal(t, byte(1), rejectPacket[ipv6.HeaderLen+1]) + // Verify entire original packet is included in body + assert.Equal(t, packet, rejectPacket[ipv6.HeaderLen+8:]) + + // Large packet: body is truncated to 1000 bytes + largePkt := makeIPv6Packet(src, dst, 17, make([]byte, 1200)) + rejectPacket = CreateRejectPacket(largePkt, out) + assert.NotNil(t, rejectPacket) + assert.Len(t, rejectPacket, ipv6.HeaderLen+8+1000) + assert.Equal(t, largePkt[:1000], rejectPacket[ipv6.HeaderLen+8:]) +} + +func Test_CreateRejectPacketIPv6_TCP(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + + // TCP SYN packet (next header 6) + tcpPayload := make([]byte, 20) + tcpPayload[0] = 0x00 // src port high + tcpPayload[1] = 0x50 // src port low (80) + tcpPayload[2] = 0x01 // dst port high + tcpPayload[3] = 0xBB // dst port low (443) + binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq + binary.BigEndian.PutUint32(tcpPayload[8:], 0) // ack seq + tcpPayload[12] = (20 >> 2) << 4 // data offset + tcpPayload[13] = 0b00000010 // SYN flag + + packet := makeIPv6Packet(src, dst, 6, tcpPayload) + + out := make([]byte, MaxRejectPacketSize) + rejectPacket := CreateRejectPacket(packet, out) + assert.NotNil(t, rejectPacket) + + // Expected: 40 (ipv6 hdr) + 20 (tcp RST) + expectedLen := ipv6.HeaderLen + 20 + assert.Len(t, rejectPacket, expectedLen) + + // Verify version + assert.Equal(t, byte(ipv6.Version<<4), rejectPacket[0]&0xf0) + // Verify next header is TCP (6) + assert.Equal(t, byte(6), rejectPacket[6]) + // Verify src/dst are swapped + assert.Equal(t, dst.To16(), net.IP(rejectPacket[8:24])) + assert.Equal(t, src.To16(), net.IP(rejectPacket[24:40])) + // Verify ports are swapped + tcpOut := rejectPacket[ipv6.HeaderLen:] + assert.Equal(t, uint16(443), binary.BigEndian.Uint16(tcpOut[0:2])) + assert.Equal(t, uint16(80), binary.BigEndian.Uint16(tcpOut[2:4])) + // RST+ACK flags (since input was SYN without ACK) + assert.Equal(t, byte(0b00010100), tcpOut[13]) + // ack_seq = original seq (1000) + SYN (1) + FIN (0) + segment data (0) + assert.Equal(t, uint32(1001), binary.BigEndian.Uint32(tcpOut[8:])) +} + +func Test_CreateRejectPacketIPv6_TCPWithACK(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + + // TCP packet with ACK set + tcpPayload := make([]byte, 20) + tcpPayload[0] = 0x00 + tcpPayload[1] = 0x50 + tcpPayload[2] = 0x01 + tcpPayload[3] = 0xBB + binary.BigEndian.PutUint32(tcpPayload[4:], 1000) // seq + binary.BigEndian.PutUint32(tcpPayload[8:], 2000) // ack seq + tcpPayload[12] = (20 >> 2) << 4 // data offset + tcpPayload[13] = 0b00010000 // ACK flag + + packet := makeIPv6Packet(src, dst, 6, tcpPayload) + + out := make([]byte, MaxRejectPacketSize) + rejectPacket := CreateRejectPacket(packet, out) + assert.NotNil(t, rejectPacket) + + tcpOut := rejectPacket[ipv6.HeaderLen:] + // RST only (no ACK) since input had ACK + assert.Equal(t, byte(0b00000100), tcpOut[13]) + // seq = original ack_seq + assert.Equal(t, uint32(2000), binary.BigEndian.Uint32(tcpOut[4:])) +} + +func Test_CreateRejectPacketIPv6_NoICMPError(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + out := make([]byte, MaxRejectPacketSize) + + // ICMPv6 error types (1-4) should not generate reject packets + for icmpType := byte(1); icmpType <= 4; icmpType++ { + payload := make([]byte, 8) + payload[0] = icmpType + packet := makeIPv6Packet(src, dst, 58, payload) + + rejectPacket := CreateRejectPacket(packet, out) + assert.Nil(t, rejectPacket, "ICMPv6 type %d should not generate a reject packet", icmpType) + } + + // ICMPv6 non-error types should still generate reject packets + nonErrorTypes := []byte{128, 129, 133, 134} + for _, icmpType := range nonErrorTypes { + payload := make([]byte, 8) + payload[0] = icmpType + packet := makeIPv6Packet(src, dst, 58, payload) + + rejectPacket := CreateRejectPacket(packet, out) + assert.NotNil(t, rejectPacket, "ICMPv6 type %d should generate a reject packet", icmpType) + } +} + +func Test_CreateRejectPacketIPv6_TooShort(t *testing.T) { + // Packet too short to be valid IPv6 + out := make([]byte, MaxRejectPacketSize) + assert.Nil(t, CreateRejectPacket([]byte{0x60}, out)) + assert.Nil(t, CreateRejectPacket(make([]byte, 39), out)) +} + +func Test_CreateRejectPacketIPv6_ExtensionHeaders(t *testing.T) { + src := net.ParseIP("fd00::1") + dst := net.ParseIP("fd00::2") + + // IPv6 + Hop-by-Hop extension header + TCP + hopByHop := []byte{ + 6, // next header: TCP + 0, // length (8 bytes total) + 0, 0, // padding + 0, 0, 0, 0, + } + tcpPayload := make([]byte, 20) + tcpPayload[0] = 0x00 + tcpPayload[1] = 0x50 + tcpPayload[2] = 0x01 + tcpPayload[3] = 0xBB + binary.BigEndian.PutUint32(tcpPayload[4:], 1000) + binary.BigEndian.PutUint32(tcpPayload[8:], 2000) + tcpPayload[12] = (20 >> 2) << 4 + tcpPayload[13] = 0b00010000 // ACK + + payload := append(hopByHop, tcpPayload...) + packet := makeIPv6Packet(src, dst, 0, payload) // next header 0 = Hop-by-Hop + + out := make([]byte, MaxRejectPacketSize) + rejectPacket := CreateRejectPacket(packet, out) + assert.NotNil(t, rejectPacket) + + // Should produce TCP RST + expectedLen := ipv6.HeaderLen + 20 + assert.Len(t, rejectPacket, expectedLen) + assert.Equal(t, byte(6), rejectPacket[6]) // next header is TCP + tcpOut := rejectPacket[ipv6.HeaderLen:] + assert.Equal(t, byte(0b00000100), tcpOut[13]) // RST only +} + +func TestCreateICMPEchoResponse_IPv4(t *testing.T) { + // Build a simple IPv4 ICMP Echo Request + packet := make([]byte, 28) + packet[0] = 0x45 // version 4, IHL 5 + binary.BigEndian.PutUint16(packet[2:], uint16(28)) // total length + packet[8] = 64 // TTL + packet[9] = 1 // protocol ICMP + copy(packet[12:16], net.IPv4(10, 0, 0, 1).To4()) // src + copy(packet[16:20], net.IPv4(10, 0, 0, 2).To4()) // dst + packet[20] = 8 // ICMP Echo Request + + out := make([]byte, len(packet)) + result := CreateICMPEchoResponse(packet, out) + assert.NotNil(t, result) + assert.Equal(t, byte(0x45), result[0]) + // src/dst swapped + assert.Equal(t, net.IPv4(10, 0, 0, 2).To4(), net.IP(result[12:16])) + assert.Equal(t, net.IPv4(10, 0, 0, 1).To4(), net.IP(result[16:20])) + // ICMP Echo Reply + assert.Equal(t, byte(0), result[20]) +} + +func TestCreateICMPEchoResponse_IPv6(t *testing.T) { + src := net.ParseIP("fd00::1").To16() + dst := net.ParseIP("fd00::2").To16() + + // Build an IPv6 ICMPv6 Echo Request packet + // IPv6 header (40 bytes) + ICMPv6 (8 bytes) + packet := make([]byte, 48) + packet[0] = 0x60 // version 6 + payloadLen := uint16(8) // ICMPv6 header only + binary.BigEndian.PutUint16(packet[4:], payloadLen) + packet[6] = 58 // Next Header: ICMPv6 + packet[7] = 64 // Hop Limit + copy(packet[8:24], src) // src address + copy(packet[24:40], dst) // dst address + + // ICMPv6 Echo Request + icmp := packet[40:] + icmp[0] = 128 // type: Echo Request + icmp[1] = 0 // code + binary.BigEndian.PutUint16(icmp[4:], 1) // identifier + binary.BigEndian.PutUint16(icmp[6:], 1) // sequence number + + // Compute correct checksum for the request + csum := ipv6PseudoheaderChecksum(src, dst, 58, uint32(payloadLen)) + binary.BigEndian.PutUint16(icmp[2:], tcpipChecksum(icmp, csum)) + + out := make([]byte, len(packet)) + result := CreateICMPEchoResponse(packet, out) + assert.NotNil(t, result) + + // Version should still be 6 + assert.Equal(t, byte(6), result[0]>>4) + // src/dst swapped + assert.Equal(t, dst, net.IP(result[8:24])) + assert.Equal(t, src, net.IP(result[24:40])) + // ICMPv6 Echo Reply type + assert.Equal(t, byte(129), result[40]) + + // Verify checksum is valid (tcpipChecksum returns 0 when data+checksum is correct) + respIcmp := result[40:] + verifyCsum := ipv6PseudoheaderChecksum(result[8:24], result[24:40], 58, uint32(payloadLen)) + assert.Equal(t, uint16(0), tcpipChecksum(respIcmp, verifyCsum)) +} + +func TestCreateICMPEchoResponse_IPv6_NotEchoRequest(t *testing.T) { + src := net.ParseIP("fd00::1").To16() + dst := net.ParseIP("fd00::2").To16() + + packet := make([]byte, 48) + packet[0] = 0x60 + binary.BigEndian.PutUint16(packet[4:], 8) + packet[6] = 58 + packet[7] = 64 + copy(packet[8:24], src) + copy(packet[24:40], dst) + + // ICMPv6 type 1 (Destination Unreachable) - not Echo Request + packet[40] = 1 + + out := make([]byte, len(packet)) + result := CreateICMPEchoResponse(packet, out) + assert.Nil(t, result) +} + +func TestCreateICMPEchoResponse_IPv6_NotICMPv6(t *testing.T) { + src := net.ParseIP("fd00::1").To16() + dst := net.ParseIP("fd00::2").To16() + + packet := make([]byte, 48) + packet[0] = 0x60 + binary.BigEndian.PutUint16(packet[4:], 8) + packet[6] = 6 // TCP, not ICMPv6 + packet[7] = 64 + copy(packet[8:24], src) + copy(packet[24:40], dst) + + out := make([]byte, len(packet)) + result := CreateICMPEchoResponse(packet, out) + assert.Nil(t, result) +} 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/outside.go b/outside.go index 6683a750..dd5b59bf 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 @@ -181,6 +182,13 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, if err != nil { return } + // Advance the replay window now that the frame is authenticated + if !hostinfo.ConnectionState.window.Update(f.l, h.MessageCounter) { + if f.l.Enabled(context.Background(), slog.LevelDebug) { + hostinfo.logger(f.l).Debug("dropping out of window relay packet", "header", h) + } + return + } // Successfully validated the thing. Get rid of the Relay header. signedPayload = signedPayload[header.Len:] // Pull the Roaming parts up here, and return in all call paths. @@ -269,15 +277,16 @@ 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 hostinfo.multiportRx { // If the remote is sending with multiport, we aren't roaming unless // the IP has changed - if hostinfo.remote.Addr().Compare(via.UdpAddr.Addr()) == 0 { + if curRemote.Addr().Compare(via.UdpAddr.Addr()) == 0 { return } // Keep the port from the original hostinfo, because the remote is transmitting from multiport ports - via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), hostinfo.remote.Port()) + via.UdpAddr = netip.AddrPortFrom(via.UdpAddr.Addr(), curRemote.Port()) } if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) { if f.l.Enabled(context.Background(), slog.LevelDebug) { @@ -290,7 +299,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, ) } @@ -298,11 +307,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) } @@ -422,16 +431,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 { @@ -591,10 +598,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_darwin.go b/overlay/tun_darwin.go index 524ef0cd..d30148b9 100644 --- a/overlay/tun_darwin.go +++ b/overlay/tun_darwin.go @@ -23,7 +23,7 @@ import ( ) type tun struct { - io.ReadWriteCloser + f *os.File Device string vpnNetworks []netip.Prefix DefaultMTU int @@ -31,9 +31,6 @@ type tun struct { routeTree atomic.Pointer[bart.Table[routing.Gateways]] linkAddr *netroute.LinkAddr l *slog.Logger - - // cache out buffer since we need to prepend 4 bytes for tun metadata - out []byte } type ifReq struct { @@ -124,11 +121,11 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t } t := &tun{ - ReadWriteCloser: os.NewFile(uintptr(fd), ""), - Device: name, - vpnNetworks: vpnNetworks, - DefaultMTU: c.GetInt("tun.mtu", DefaultMTU), - l: l, + f: os.NewFile(uintptr(fd), ""), + Device: name, + vpnNetworks: vpnNetworks, + DefaultMTU: c.GetInt("tun.mtu", DefaultMTU), + l: l, } err = t.reload(c, true) @@ -158,8 +155,8 @@ func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, e } func (t *tun) Close() error { - if t.ReadWriteCloser != nil { - return t.ReadWriteCloser.Close() + if t.f != nil { + return t.f.Close() } return nil } @@ -502,42 +499,103 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error { return nil } +// tunWritev and tunReadv are linkname'd to x/sys/unix's libc-routed writev/readv stubs so the +// calls go through libSystem's pinned trampoline. A raw syscall.Syscall(SYS_WRITEV/SYS_READV, ...) +// on darwin/arm64 emits an SVC #0x80 trap (see $GOROOT/src/syscall/asm_darwin_arm64.s), the path +// Apple keeps warning they will eventually disallow. We pull the low-level stubs instead of calling +// unix.Writev/unix.Readv because those take [][]byte and rebuild the []Iovec every call, which +// heap-allocates the header; linkname'ing the stubs lets us hand them our own stack-allocated +// iovecs. See golang/go#78049. + +//go:linkname tunWritev golang.org/x/sys/unix.writev +//go:noescape +func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error) + +//go:linkname tunReadv golang.org/x/sys/unix.readv +//go:noescape +func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error) + +// Read pulls one IP packet off the utun device, scattering the 4 byte protocol header away from +// the packet so the payload lands directly in to. func (t *tun) Read(to []byte) (int, error) { - buf := make([]byte, len(to)+4) + var head [4]byte - n, err := t.ReadWriteCloser.Read(buf) + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - copy(to, buf[4:]) - return n - 4, err + var n int + var callErr error + err = rc.Read(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &to[0], Len: uint64(len(to))}, + } + n, callErr = tunReadv(int(fd), iovecs) + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + if n < 4 { + return 0, nil + } + return n - 4, nil } -// Write is only valid for single threaded use +// Write pushes one IP packet onto the utun device. func (t *tun) Write(from []byte) (int, error) { - buf := t.out - if cap(buf) < len(from)+4 { - buf = make([]byte, len(from)+4) - t.out = buf - } - buf = buf[:len(from)+4] - if len(from) == 0 { return 0, syscall.EIO } - // Determine the IP Family for the NULL L2 Header ipVer := from[0] >> 4 - if ipVer == 4 { - buf[3] = syscall.AF_INET - } else if ipVer == 6 { - buf[3] = syscall.AF_INET6 - } else { + var head [4]byte + switch ipVer { + case 4: + head[3] = syscall.AF_INET + case 6: + head[3] = syscall.AF_INET6 + default: return 0, fmt.Errorf("unable to determine IP version from packet") } - copy(buf[4:], from) + // Grab rc as a local so the compiler can devirtualize the call and keep the closure on the stack. + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - n, err := t.ReadWriteCloser.Write(buf) - return n - 4, err + var n int + var callErr error + err = rc.Write(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &from[0], Len: uint64(len(from))}, + } + n, callErr = tunWritev(int(fd), iovecs) + // Type-assert to syscall.Errno so the EAGAIN/EWOULDBLOCK/EINTR check doesn't box the errno + // constants into error interfaces on every call. + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + + return n - 4, nil } func (t *tun) Networks() []netip.Prefix { diff --git a/overlay/tun_openbsd.go b/overlay/tun_openbsd.go index 81362184..41224777 100644 --- a/overlay/tun_openbsd.go +++ b/overlay/tun_openbsd.go @@ -57,8 +57,6 @@ type tun struct { l *slog.Logger f *os.File fd int - // cache out buffer since we need to prepend 4 bytes for tun metadata - out []byte } var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`) @@ -124,42 +122,103 @@ func (t *tun) Close() error { return nil } +// tunWritev and tunReadv are linkname'd to x/sys/unix's libc-routed writev/readv stubs so the +// calls go through libc's pinned trampoline. OpenBSD's pinsyscall protection rejects a raw +// syscall.Syscall(SYS_WRITEV/SYS_READV, ...) because it doesn't originate from a libc-pinned +// address, so we can't use the syscall.Syscall pattern that freebsd / netbsd use. We pull the +// low-level stubs instead of calling unix.Writev/unix.Readv because those take [][]byte and rebuild +// the []Iovec every call, which heap-allocates the header; linkname'ing the stubs lets us hand them +// our own stack-allocated iovecs. See golang/go#78049. + +//go:linkname tunWritev golang.org/x/sys/unix.writev +//go:noescape +func tunWritev(fd int, iovecs []unix.Iovec) (n int, err error) + +//go:linkname tunReadv golang.org/x/sys/unix.readv +//go:noescape +func tunReadv(fd int, iovecs []unix.Iovec) (n int, err error) + +// Read pulls one IP packet off the tun device, scattering the 4 byte protocol header away from the +// packet so the payload lands directly in to. func (t *tun) Read(to []byte) (int, error) { - buf := make([]byte, len(to)+4) + var head [4]byte - n, err := t.f.Read(buf) + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - copy(to, buf[4:]) - return n - 4, err + var n int + var callErr error + err = rc.Read(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &to[0], Len: uint64(len(to))}, + } + n, callErr = tunReadv(int(fd), iovecs) + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + if n < 4 { + return 0, nil + } + return n - 4, nil } -// Write is only valid for single threaded use +// Write pushes one IP packet onto the tun device. func (t *tun) Write(from []byte) (int, error) { - buf := t.out - if cap(buf) < len(from)+4 { - buf = make([]byte, len(from)+4) - t.out = buf - } - buf = buf[:len(from)+4] - if len(from) == 0 { return 0, syscall.EIO } - // Determine the IP Family for the NULL L2 Header ipVer := from[0] >> 4 - if ipVer == 4 { - buf[3] = syscall.AF_INET - } else if ipVer == 6 { - buf[3] = syscall.AF_INET6 - } else { + var head [4]byte + switch ipVer { + case 4: + head[3] = syscall.AF_INET + case 6: + head[3] = syscall.AF_INET6 + default: return 0, fmt.Errorf("unable to determine IP version from packet") } - copy(buf[4:], from) + // Grab rc as a local so the compiler can devirtualize the call and keep the closure on the stack. + rc, err := t.f.SyscallConn() + if err != nil { + return 0, err + } - n, err := t.f.Write(buf) - return n - 4, err + var n int + var callErr error + err = rc.Write(func(fd uintptr) bool { + iovecs := []unix.Iovec{ + {Base: &head[0], Len: 4}, + {Base: &from[0], Len: uint64(len(from))}, + } + n, callErr = tunWritev(int(fd), iovecs) + // Type-assert to syscall.Errno so the EAGAIN/EWOULDBLOCK/EINTR check doesn't box the errno + // constants into error interfaces on every call. + if errno, ok := callErr.(syscall.Errno); ok && errno.Temporary() { + return false + } + return true + }) + if err != nil { + return 0, err + } + if callErr != nil { + return 0, callErr + } + + return n - 4, nil } func (t *tun) addIp(cidr netip.Prefix) error { 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 1fd98963..318a9f1a 100644 --- a/relay_manager.go +++ b/relay_manager.go @@ -94,7 +94,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak } relayHostInfo := rm.hostmap.QueryVpnAddr(relay) - if relayHostInfo == nil || !relayHostInfo.remote.IsValid() { + if relayHostInfo == nil || !relayHostInfo.GetRemote().IsValid() { hl.Log(context.Background(), level, "Establish tunnel to relay target", "relay", relay.String()) f.Handshake(relay) continue @@ -104,7 +104,7 @@ func (rm *relayManager) StartRelays(f *Interface, vpnIp netip.Addr, hh *Handshak existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp) if !ok { // No relays exist or requested yet. - if relayHostInfo.remote.IsValid() { + if relayHostInfo.GetRemote().IsValid() { idx, err := AddRelay(rm.l, relayHostInfo, rm.hostmap, vpnIp, nil, TerminalType, Requested) if err != nil { hl.Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err) @@ -309,6 +309,22 @@ func (rm *relayManager) HandleControlMsg(h *HostInfo, d []byte, f *Interface) { v = cert.Version2 } + // validate: + switch msg.Type { + case NebulaControl_CreateRelayRequest, NebulaControl_CreateRelayResponse: + if msg.RelayFromAddr == nil { + if f.l.Enabled(context.Background(), slog.LevelDebug) { + h.logger(f.l).Debug("Control message received with nil RelayFromAddr", "type", msg.Type) + } + return + } else if msg.RelayToAddr == nil { + if f.l.Enabled(context.Background(), slog.LevelDebug) { + h.logger(f.l).Debug("Control message received with nil RelayToAddr", "type", msg.Type) + } + return + } + } + switch msg.Type { case NebulaControl_CreateRelayRequest: rm.handleCreateRelayRequest(v, h, f, msg) @@ -318,17 +334,17 @@ 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", - "relayFrom", protoAddrToNetAddr(m.RelayFromAddr), - "relayTo", protoAddrToNetAddr(m.RelayToAddr), + "relayFrom", relayFrom, + "relayTo", relayTo, "initiatorRelayIndex", m.InitiatorRelayIndex, "responderRelayIndex", m.ResponderRelayIndex, "vpnAddrs", h.vpnAddrs, ) - target := m.RelayToAddr - targetAddr := protoAddrToNetAddr(target) - relay, err := rm.EstablishRelay(h, m) if err != nil { rm.l.Error("Failed to update relay for relayTo", "error", err) @@ -344,7 +360,7 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f rm.l.Error("Can't find a HostInfo for peer", "relayTo", relay.PeerAddr) return } - peerRelay, ok := peerHostInfo.relayState.QueryRelayForByIp(targetAddr) + peerRelay, ok := peerHostInfo.relayState.QueryRelayForByIp(relayTo) if !ok { rm.l.Error("peerRelay does not have Relay state for relayTo", "relayTo", peerHostInfo.vpnAddrs[0]) return @@ -354,19 +370,19 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f // I initiated the request to this peer, but haven't heard back from the peer yet. I must wait for this peer // to respond to complete the connection. case PeerRequested, Disestablished, Established: - peerHostInfo.relayState.UpdateRelayForByIpState(targetAddr, Established) + peerHostInfo.relayState.UpdateRelayForByIpState(relayTo, Established) resp := NebulaControl{ Type: NebulaControl_CreateRelayResponse, ResponderRelayIndex: peerRelay.LocalIndex, InitiatorRelayIndex: peerRelay.RemoteIndex, } + peer := peerHostInfo.vpnAddrs[0] if v == cert.Version1 { - peer := peerHostInfo.vpnAddrs[0] if !peer.Is4() { rm.l.Error("Refusing to CreateRelayResponse for a v1 relay with an ipv6 address", "relayFrom", peer, - "relayTo", target, + "relayTo", relayTo, "initiatorRelayIndex", resp.InitiatorRelayIndex, "responderRelayIndex", resp.ResponderRelayIndex, "vpnAddrs", peerHostInfo.vpnAddrs, @@ -376,30 +392,31 @@ func (rm *relayManager) handleCreateRelayResponse(v cert.Version, h *HostInfo, f b := peer.As4() resp.OldRelayFromAddr = binary.BigEndian.Uint32(b[:]) - b = targetAddr.As4() + b = relayTo.As4() resp.OldRelayToAddr = binary.BigEndian.Uint32(b[:]) } else { - resp.RelayFromAddr = netAddrToProtoAddr(peerHostInfo.vpnAddrs[0]) - resp.RelayToAddr = target + resp.RelayFromAddr = netAddrToProtoAddr(peer) + resp.RelayToAddr = m.RelayToAddr } msg, err := resp.Marshal() if err != nil { rm.l.Error("relayManager Failed to marshal Control CreateRelayResponse message to create relay", "error", err) - } else { - f.SendMessageToHostInfo(header.Control, 0, peerHostInfo, msg, make([]byte, 12), make([]byte, mtu)) - rm.l.Info("send CreateRelayResponse", - "relayFrom", resp.RelayFromAddr, - "relayTo", resp.RelayToAddr, - "initiatorRelayIndex", resp.InitiatorRelayIndex, - "responderRelayIndex", resp.ResponderRelayIndex, - "vpnAddrs", peerHostInfo.vpnAddrs, - ) + return } + f.SendMessageToHostInfo(header.Control, 0, peerHostInfo, msg, make([]byte, 12), make([]byte, mtu)) + rm.l.Info("send CreateRelayResponse", + "relayFrom", peer, + "relayTo", relayTo, + "initiatorRelayIndex", resp.InitiatorRelayIndex, + "responderRelayIndex", resp.ResponderRelayIndex, + "vpnAddrs", peerHostInfo.vpnAddrs, + ) } } 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) @@ -509,7 +526,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/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/stats.go b/stats.go index 97ce7cf5..98c2baf2 100644 --- a/stats.go +++ b/stats.go @@ -331,7 +331,7 @@ func loadStatsConfig(c *config.C) (statsConfig, error) { } cfg.interval = c.GetDuration("stats.interval", 0) - if cfg.interval == 0 { + if cfg.interval <= 0 { return cfg, fmt.Errorf("stats.interval was an invalid duration: %s", c.GetString("stats.interval", "")) } diff --git a/udp/udp_darwin.go b/udp/udp_darwin.go index 8a4f5b18..3d6b39a5 100644 --- a/udp/udp_darwin.go +++ b/udp/udp_darwin.go @@ -175,8 +175,8 @@ func (u *StdConn) ListenOut(r EncReader) error { if errors.Is(err, net.ErrClosed) { return err } - u.l.Error("unexpected udp socket receive error", "error", err) + continue } r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])