Compare commits

..

16 Commits

Author SHA1 Message Date
JackDoan f5ddff5ca1 fix after rebase 2026-05-11 11:14:25 -05:00
JackDoan 400cbc26a1 parse? 2026-05-11 11:14:25 -05:00
JackDoan 01b31360df SPICY 2026-05-11 11:14:25 -05:00
JackDoan 5bdf645b0b checkpt, try to parse packets only once pt2 2026-05-11 11:14:25 -05:00
JackDoan 0375aff451 checkpt, try to parse packets only once 2026-05-11 11:14:25 -05:00
JackDoan 6cb00c613c faster
grr heap usage!
2026-05-11 11:14:25 -05:00
JackDoan 40b4ae7fb4 no 2026-05-11 11:14:25 -05:00
JackDoan cf51b6dfd7 use clear() 2026-05-11 11:14:25 -05:00
JackDoan fe93ebd017 remove udp-level RX reorder buf 2026-05-11 11:14:25 -05:00
JackDoan 961ddbfbc1 make relays take the fast path maybe 2026-05-11 11:14:25 -05:00
JackDoan 67bd9e848a scoot pinning around 2026-05-11 11:14:25 -05:00
JackDoan bc3f5d0400 scoot stuff around for e2e 2026-05-11 11:14:25 -05:00
JackDoan aef8e39cc4 disable sort-on-RX, CPU pinning seems to work for now 2026-05-11 11:14:25 -05:00
JackDoan 69863d6c81 switch to ASM vector checksum 2026-05-11 11:14:25 -05:00
JackDoan 5d35351437 GSO/GRO offloads, with TCP+ECN and UDP support 2026-05-11 11:14:25 -05:00
JackDoan f95857b4c3 better and batched tun interface 2026-05-11 11:09:10 -05:00
132 changed files with 3073 additions and 7888 deletions
+2 -5
View File
@@ -25,9 +25,9 @@ inputs:
required: false required: false
default: "code-signer" default: "code-signer"
key-prefix: key-prefix:
description: "S3 key prefix to write under; defaults to code-signing/<owner>/<repo> of the calling repo" description: "S3 key prefix the caller is authorized to write under"
required: false required: false
default: "" default: "code-signing/slackhq/nebula"
runs: runs:
using: composite using: composite
@@ -57,9 +57,6 @@ runs:
KEY_PREFIX: ${{ inputs.key-prefix }} KEY_PREFIX: ${{ inputs.key-prefix }}
run: | run: |
set -eu set -eu
# Default the prefix to this repo so the S3 key attributes the sign correctly.
# nebula-nightly runs this same action but writes under its own repo's prefix.
KEY_PREFIX="${KEY_PREFIX:-code-signing/$GITHUB_REPOSITORY}"
RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}" RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}"
find "$SIGN_PATH" -name '*.exe' -print | while read -r path find "$SIGN_PATH" -name '*.exe' -print | while read -r path
+34
View File
@@ -0,0 +1,34 @@
name: gofmt
on:
push:
branches:
- master
pull_request:
paths:
- '.github/workflows/gofmt.yml'
- '**.go'
jobs:
gofmt:
name: Run gofmt
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v6
- uses: actions/setup-go@v6
with:
go-version: '1.25'
check-latest: true
- name: Install goimports
run: |
go install golang.org/x/tools/cmd/goimports@latest
- name: gofmt
run: |
if [ "$(find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -l)" ]
then
find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -d
exit 1
fi
+6 -9
View File
@@ -10,7 +10,7 @@ jobs:
name: Build Linux/BSD All name: Build Linux/BSD All
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -36,7 +36,7 @@ jobs:
id-token: write id-token: write
contents: read contents: read
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -76,7 +76,7 @@ jobs:
HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }} HAS_SIGNING_CREDS: ${{ secrets.AC_USERNAME != '' }}
runs-on: macos-latest runs-on: macos-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -134,7 +134,7 @@ jobs:
# be overwritten # be overwritten
- name: Checkout code - name: Checkout code
if: ${{ env.HAS_DOCKER_CREDS == 'true' }} if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
uses: actions/checkout@v7 uses: actions/checkout@v6
- name: Download artifacts - name: Download artifacts
if: ${{ env.HAS_DOCKER_CREDS == 'true' }} if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
@@ -163,17 +163,14 @@ jobs:
mkdir -p build/linux-{amd64,arm64} mkdir -p build/linux-{amd64,arm64}
tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/ tar -zxvf artifacts/nebula-linux-amd64.tar.gz -C build/linux-amd64/
tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/ tar -zxvf artifacts/nebula-linux-arm64.tar.gz -C build/linux-arm64/
docker buildx build . --push -f docker/Dockerfile --platform linux/amd64,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}"
--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: release:
name: Create and Upload Release name: Create and Upload Release
needs: [build-linux, build-darwin, build-windows] needs: [build-linux, build-darwin, build-windows]
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- name: Download artifacts - name: Download artifacts
uses: actions/download-artifact@v8 uses: actions/download-artifact@v8
+3 -3
View File
@@ -30,7 +30,7 @@ jobs:
VAGRANT_DEFAULT_PROVIDER: libvirt VAGRANT_DEFAULT_PROVIDER: libvirt
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -62,7 +62,7 @@ jobs:
VAGRANT_DEFAULT_PROVIDER: virtualbox VAGRANT_DEFAULT_PROVIDER: virtualbox
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -88,7 +88,7 @@ jobs:
runs-on: windows-latest runs-on: windows-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
+1 -9
View File
@@ -18,7 +18,7 @@ jobs:
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -36,14 +36,6 @@ jobs:
working-directory: ./.github/workflows/smoke working-directory: ./.github/workflows/smoke
run: ./smoke.sh 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 - name: setup relay docker image
working-directory: ./.github/workflows/smoke working-directory: ./.github/workflows/smoke
run: ./build-relay.sh run: ./build-relay.sh
+7 -20
View File
@@ -5,19 +5,6 @@ set -e -x
rm -rf ./build rm -rf ./build
mkdir ./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 # 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 # 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 # placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
@@ -44,24 +31,24 @@ LIGHTHOUSE_IP="203.0.113.2"
../genconfig.sh >lighthouse1.yml ../genconfig.sh >lighthouse1.yml
HOST="host2" \ HOST="host2" \
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \ LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
../genconfig.sh >host2.yml ../genconfig.sh >host2.yml
HOST="host3" \ HOST="host3" \
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \ LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \ INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
../genconfig.sh >host3.yml ../genconfig.sh >host3.yml
HOST="host4" \ HOST="host4" \
LIGHTHOUSES="$LIGHTHOUSE_NIP $LIGHTHOUSE_IP:4242" \ LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \ OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
../genconfig.sh >host4.yml ../genconfig.sh >host4.yml
../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test" ../../../../nebula-cert ca -curve "${CURVE:-25519}" -name "Smoke Test"
../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "$LIGHTHOUSE_NIP/24" ../../../../nebula-cert sign -name "lighthouse1" -groups "lighthouse,lighthouse1" -ip "192.168.100.1/24"
../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "$HOST2_NIP/24" ../../../../nebula-cert sign -name "host2" -groups "host,host2" -ip "192.168.100.2/24"
../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "$HOST3_NIP/24" ../../../../nebula-cert sign -name "host3" -groups "host,host3" -ip "192.168.100.3/24"
../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "$HOST4_NIP/24" ../../../../nebula-cert sign -name "host4" -groups "host,host4" -ip "192.168.100.4/24"
) )
docker build -t "nebula:${NAME:-smoke}" . docker build -t "nebula:${NAME:-smoke}" .
+23 -36
View File
@@ -47,19 +47,6 @@ HOST2_IP="$PREFIX.3"
HOST3_IP="$PREFIX.4" HOST3_IP="$PREFIX.4"
HOST4_IP="$PREFIX.5" 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. # 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. # 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 for f in build/host2.yml build/host3.yml build/host4.yml; do
@@ -93,28 +80,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 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 host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host4.outside.pcap &
docker exec host2 ncat -nklv 2000 & docker exec host2 ncat -nklv 0.0.0.0 2000 &
docker exec host3 ncat -nklv 2000 & docker exec host3 ncat -nklv 0.0.0.0 2000 &
docker exec host4 ncat -e '/usr/bin/echo helloagainfromhost4' -nkluv 4000 & 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 3000 & 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 3000 & docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 &
set +x set +x
echo echo
echo " *** Testing ping from lighthouse1" echo " *** Testing ping from lighthouse1"
echo echo
set -x set -x
docker exec lighthouse1 ping -c1 $HOST2_NIP docker exec lighthouse1 ping -c1 192.168.100.2
docker exec lighthouse1 ping -c1 $HOST3_NIP docker exec lighthouse1 ping -c1 192.168.100.3
set +x set +x
echo echo
echo " *** Testing ping from host2" echo " *** Testing ping from host2"
echo echo
set -x set -x
docker exec host2 ping -c1 $LIGHTHOUSE_NIP docker exec host2 ping -c1 192.168.100.1
# Should fail because not allowed by host3 inbound firewall # Should fail because not allowed by host3 inbound firewall
! docker exec host2 ping -c1 $HOST3_NIP -w5 || exit 1 ! docker exec host2 ping -c1 192.168.100.3 -w5 || exit 1
set +x set +x
echo echo
@@ -122,34 +109,34 @@ echo " *** Testing ncat from host2"
echo echo
set -x set -x
# Should fail because not allowed by host3 inbound firewall # Should fail because not allowed by host3 inbound firewall
! docker exec host2 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1 ! docker exec host2 ncat -nzv -w5 192.168.100.3 2000 || exit 1
! docker exec host2 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1 ! docker exec host2 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
set +x set +x
echo echo
echo " *** Testing ping from host3" echo " *** Testing ping from host3"
echo echo
set -x set -x
docker exec host3 ping -c1 $LIGHTHOUSE_NIP docker exec host3 ping -c1 192.168.100.1
docker exec host3 ping -c1 $HOST2_NIP docker exec host3 ping -c1 192.168.100.2
set +x set +x
echo echo
echo " *** Testing ncat from host3" echo " *** Testing ncat from host3"
echo echo
set -x set -x
docker exec host3 ncat -nzv -w5 $HOST2_NIP 2000 docker exec host3 ncat -nzv -w5 192.168.100.2 2000
docker exec host3 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2 docker exec host3 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2
set +x set +x
echo echo
echo " *** Testing ping from host4" echo " *** Testing ping from host4"
echo echo
set -x set -x
docker exec host4 ping -c1 $LIGHTHOUSE_NIP docker exec host4 ping -c1 192.168.100.1
# Should fail because not allowed by host4 outbound firewall # Should fail because not allowed by host4 outbound firewall
! docker exec host4 ping -c1 $HOST2_NIP -w5 || exit 1 ! docker exec host4 ping -c1 192.168.100.2 -w5 || exit 1
! docker exec host4 ping -c1 $HOST3_NIP -w5 || exit 1 ! docker exec host4 ping -c1 192.168.100.3 -w5 || exit 1
set +x set +x
echo echo
@@ -157,10 +144,10 @@ echo " *** Testing ncat from host4"
echo echo
set -x set -x
# Should fail because not allowed by host4 outbound firewall # Should fail because not allowed by host4 outbound firewall
! docker exec host4 ncat -nzv -w5 $HOST2_NIP 2000 || exit 1 ! docker exec host4 ncat -nzv -w5 192.168.100.2 2000 || exit 1
! docker exec host4 ncat -nzv -w5 $HOST3_NIP 2000 || exit 1 ! docker exec host4 ncat -nzv -w5 192.168.100.3 2000 || exit 1
! docker exec host4 ncat -nzuv -w5 $HOST2_NIP 3000 | grep -q host2 || exit 1 ! docker exec host4 ncat -nzuv -w5 192.168.100.2 3000 | grep -q host2 || exit 1
! docker exec host4 ncat -nzuv -w5 $HOST3_NIP 3000 | grep -q host3 || exit 1 ! docker exec host4 ncat -nzuv -w5 192.168.100.3 3000 | grep -q host3 || exit 1
set +x set +x
echo echo
@@ -172,7 +159,7 @@ set -x
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000, # 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, # conntrack must let host4's listener reply on that flow. If it doesn't,
# the echo back from host4 never reaches host2. # the echo back from host4 never reaches host2.
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv $HOST4_NIP 4000" | grep -q helloagainfromhost4 docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4
docker exec host4 sh -c 'kill 1' docker exec host4 sh -c 'kill 1'
docker exec host3 sh -c 'kill 1' docker exec host3 sh -c 'kill 1'
+72 -90
View File
@@ -13,28 +13,20 @@ on:
- 'go.sum' - 'go.sum'
jobs: jobs:
static: test-linux:
name: Static checks name: Build all and test on ubuntu-linux
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
go-version: '1.25' go-version: '1.25'
check-latest: true check-latest: true
- name: Install goimports - name: Build
run: go install golang.org/x/tools/cmd/goimports@latest run: make all
- name: gofmt
run: |
if [ "$(find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -l)" ]
then
find . -iname '*.go' | grep -v '\.pb\.go$' | xargs goimports -d
exit 1
fi
- name: Vet - name: Vet
run: make vet run: make vet
@@ -44,41 +36,27 @@ jobs:
with: with:
version: v2.5 version: v2.5
test: - name: Test
name: Test ${{ matrix.name }} run: make test
runs-on: ${{ matrix.os }}
strategy: - name: End 2 end
fail-fast: false run: make e2evv
matrix:
include: - name: Build test mobile
- name: linux run: make build-test-mobile
os: ubuntu-latest
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert - uses: actions/upload-artifact@v7
test-cmd: make test with:
e2e-cmd: make e2evv name: e2e packet flow linux-latest
- name: linux-boringcrypto path: e2e/mermaid/linux-latest
os: ubuntu-latest if-no-files-found: warn
build-cmd: make bin-boringcrypto
test-cmd: make test-boringcrypto test-linux-boringcrypto:
e2e-cmd: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0" name: Build and test on linux with boringcrypto
- name: linux-pkcs11 runs-on: ubuntu-latest
os: ubuntu-latest
build-cmd: make bin-pkcs11
test-cmd: make test-pkcs11
e2e-cmd: ''
- name: macos
os: macos-latest
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
test-cmd: make test
e2e-cmd: make e2evv
- name: windows
os: windows-latest
build-cmd: go build ./cmd/nebula ./cmd/nebula-cert
test-cmd: make test
e2e-cmd: make e2evv
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
@@ -86,65 +64,69 @@ jobs:
check-latest: true check-latest: true
- name: Build - name: Build
run: ${{ matrix.build-cmd }} run: make bin-boringcrypto
- name: Cross-build darwin-amd64
if: matrix.name == 'macos'
run: GOARCH=amd64 go build -o /tmp/nebula-amd64 ./cmd/nebula && GOARCH=amd64 go build -o /tmp/nebula-cert-amd64 ./cmd/nebula-cert
- name: Test - name: Test
run: ${{ matrix.test-cmd }} run: make test-boringcrypto
- name: End 2 end - name: End 2 end
if: matrix.e2e-cmd != '' run: make e2e GOEXPERIMENT=boringcrypto CGO_ENABLED=1 TEST_ENV="TEST_LOGS=1" TEST_FLAGS="-v -ldflags -checklinkname=0"
run: ${{ matrix.e2e-cmd }}
- uses: actions/upload-artifact@v7 test-linux-pkcs11:
if: matrix.e2e-cmd != '' && always() name: Build and test on linux with pkcs11
with:
name: e2e packet flow ${{ matrix.name }}
path: e2e/mermaid/
if-no-files-found: warn
cross-build:
name: Cross-build ${{ matrix.name }}
runs-on: ubuntu-latest runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
include:
- {name: linux-arm, make-target: all-cross-linux-arm}
- {name: linux-mips, make-target: all-cross-linux-mips}
- {name: linux-other, make-target: all-cross-linux-other}
- {name: freebsd, make-target: all-freebsd}
- {name: openbsd, make-target: all-openbsd}
- {name: netbsd, make-target: all-netbsd}
- {name: windows, make-target: all-cross-windows}
- {name: mobile, make-target: build-test-mobile}
steps: steps:
- uses: actions/checkout@v7 - uses: actions/checkout@v6
- uses: actions/setup-go@v6 - uses: actions/setup-go@v6
with: with:
go-version: '1.25' go-version: '1.25'
check-latest: true check-latest: true
- name: Build ${{ matrix.name }} - name: Build
run: make -j"$(nproc)" ${{ matrix.make-target }} run: make bin-pkcs11
finish: - name: Test
name: CI status run: make test-pkcs11
if: always()
needs: [static, test, cross-build] test:
runs-on: ubuntu-latest name: Build and test on ${{ matrix.os }}
runs-on: ${{ matrix.os }}
strategy:
matrix:
os: [windows-latest, macos-latest]
steps: steps:
- name: Fail if any upstream job failed - uses: actions/checkout@v6
if: contains(needs.*.result, 'failure') || contains(needs.*.result, 'cancelled')
run: |
echo "upstream results: ${{ toJSON(needs) }}"
exit 1
- name: All upstream jobs passed - uses: actions/setup-go@v6
run: echo "ok" with:
go-version: '1.25'
check-latest: true
- name: Build nebula
run: go build ./cmd/nebula
- name: Build nebula-cert
run: go build ./cmd/nebula-cert
- name: Vet
run: make vet
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
version: v2.5
- name: Test
run: make test
- name: End 2 end
run: make e2evv
- uses: actions/upload-artifact@v7
with:
name: e2e packet flow ${{ matrix.os }}
path: e2e/mermaid/${{ matrix.os }}
if-no-files-found: warn
+1 -49
View File
@@ -60,18 +60,6 @@ ALL = $(ALL_LINUX) \
windows-amd64 \ windows-amd64 \
windows-arm64 windows-arm64
# Cross-build shards used by .github/workflows/test.yml — same as ALL_*
# but with the arch that has a native CI runner removed, so the cross-build
# job is not duplicating coverage the native test jobs already give.
ALL_CROSS_LINUX = $(filter-out linux-amd64,$(ALL_LINUX))
# ALL_CROSS_LINUX further split into family sub-shards so each can run on
# its own CI runner in parallel. Union of the three must equal
# ALL_CROSS_LINUX; adding a new linux arch goes into the matching family.
ALL_CROSS_LINUX_ARM = linux-arm-5 linux-arm-6 linux-arm-7 linux-arm64
ALL_CROSS_LINUX_MIPS = linux-mips linux-mipsle linux-mips64 linux-mips64le linux-mips-softfloat
ALL_CROSS_LINUX_OTHER = linux-386 linux-ppc64le linux-riscv64 linux-loong64
e2e: e2e:
$(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e $(TEST_ENV) go test -tags=e2e_testing -count=1 $(TEST_FLAGS) ./e2e
@@ -94,35 +82,6 @@ DOCKER_BIN = build/linux-amd64/nebula build/linux-amd64/nebula-cert
all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert) all: $(ALL:%=build/%/nebula) $(ALL:%=build/%/nebula-cert)
all-linux: $(ALL_LINUX:%=build/%/nebula) $(ALL_LINUX:%=build/%/nebula-cert)
all-freebsd: $(ALL_FREEBSD:%=build/%/nebula) $(ALL_FREEBSD:%=build/%/nebula-cert)
all-openbsd: $(ALL_OPENBSD:%=build/%/nebula) $(ALL_OPENBSD:%=build/%/nebula-cert)
all-netbsd: $(ALL_NETBSD:%=build/%/nebula) $(ALL_NETBSD:%=build/%/nebula-cert)
all-darwin: build/darwin-amd64/nebula build/darwin-amd64/nebula-cert build/darwin-arm64/nebula build/darwin-arm64/nebula-cert
all-windows: build/windows-amd64/nebula.exe build/windows-amd64/nebula-cert.exe build/windows-arm64/nebula.exe build/windows-arm64/nebula-cert.exe
# CI cross-build shards. darwin-arm64 is covered by the native macos-latest
# job; windows-amd64 is covered by the native windows-latest job; both are
# omitted here to avoid building them a second time. darwin-amd64 stays in
# all-cross-darwin because intel mac is only a labeled/master-time native
# job, so PRs still need cross-build coverage for it.
all-cross-linux: $(ALL_CROSS_LINUX:%=build/%/nebula) $(ALL_CROSS_LINUX:%=build/%/nebula-cert)
all-cross-linux-arm: $(ALL_CROSS_LINUX_ARM:%=build/%/nebula) $(ALL_CROSS_LINUX_ARM:%=build/%/nebula-cert)
all-cross-linux-mips: $(ALL_CROSS_LINUX_MIPS:%=build/%/nebula) $(ALL_CROSS_LINUX_MIPS:%=build/%/nebula-cert)
all-cross-linux-other: $(ALL_CROSS_LINUX_OTHER:%=build/%/nebula) $(ALL_CROSS_LINUX_OTHER:%=build/%/nebula-cert)
all-cross-darwin: build/darwin-amd64/nebula build/darwin-amd64/nebula-cert
all-cross-windows: build/windows-arm64/nebula.exe build/windows-arm64/nebula-cert.exe
docker: docker/linux-$(shell go env GOARCH) docker: docker/linux-$(shell go env GOARCH)
release: $(ALL:%=build/nebula-%.tar.gz) release: $(ALL:%=build/nebula-%.tar.gz)
@@ -161,10 +120,6 @@ bin-pkcs11: BUILD_ARGS += -tags pkcs11
bin-pkcs11: CGO_ENABLED = 1 bin-pkcs11: CGO_ENABLED = 1
bin-pkcs11: bin bin-pkcs11: bin
# Build with the pprof debug server (serves on :6060). See startPprofServer.
debug: BUILD_ARGS += -tags debug
debug: bin
bin: bin:
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH} go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula${NEBULA_CMD_SUFFIX} ${NEBULA_CMD_PATH}
go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert go build $(BUILD_ARGS) -ldflags "$(LDFLAGS)" -o ./nebula-cert${NEBULA_CMD_SUFFIX} ./cmd/nebula-cert
@@ -272,9 +227,6 @@ smoke-relay-docker: bin-docker
cd .github/workflows/smoke/ && ./build-relay.sh cd .github/workflows/smoke/ && ./build-relay.sh
cd .github/workflows/smoke/ && ./smoke-relay.sh cd .github/workflows/smoke/ && ./smoke-relay.sh
smoke-docker-ipv6: export SMOKE_OVERLAY_IPV6 = 1
smoke-docker-ipv6: smoke-docker
smoke-docker-race: BUILD_ARGS = -race smoke-docker-race: BUILD_ARGS = -race
smoke-docker-race: CGO_ENABLED = 1 smoke-docker-race: CGO_ENABLED = 1
smoke-docker-race: smoke-docker smoke-docker-race: smoke-docker
@@ -284,5 +236,5 @@ smoke-vagrant/%: bin-docker build/%/nebula
cd .github/workflows/smoke/ && ./smoke-vagrant.sh $* cd .github/workflows/smoke/ && ./smoke-vagrant.sh $*
.FORCE: .FORCE:
.PHONY: all all-linux all-freebsd all-openbsd all-netbsd all-darwin all-windows all-cross-linux all-cross-linux-arm all-cross-linux-mips all-cross-linux-other all-cross-darwin all-cross-windows bench bench-cpu bench-cpu-long bin debug build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/% .PHONY: bench bench-cpu bench-cpu-long bin build-test-mobile e2e e2ev e2evv e2evvv e2evvvv proto release service smoke-docker smoke-docker-race test test-cov-html smoke-vagrant/%
.DEFAULT_GOAL := bin .DEFAULT_GOAL := bin
+4 -10
View File
@@ -13,12 +13,6 @@ import (
"golang.org/x/crypto/ed25519" "golang.org/x/crypto/ed25519"
) )
// testCertNow is the reference "now" used to derive default before/after times
// in NewTestCaCert and NewTestCert. Holding it fixed for the lifetime of the
// test binary keeps CA and leaf defaults aligned at the same second, so a leaf
// signed with default times can never expire after its CA on a rounding race.
var testCertNow = time.Now().Round(time.Second)
// NewTestCaCert will create a new ca certificate // NewTestCaCert will create a new ca certificate
func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) { func NewTestCaCert(version Version, curve Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
var err error var err error
@@ -40,10 +34,10 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
} }
if before.IsZero() { if before.IsZero() {
before = testCertNow.Add(time.Second * -60) before = time.Now().Add(time.Second * -60).Round(time.Second)
} }
if after.IsZero() { if after.IsZero() {
after = testCertNow.Add(time.Second * 60) after = time.Now().Add(time.Second * 60).Round(time.Second)
} }
t := &TBSCertificate{ t := &TBSCertificate{
@@ -76,11 +70,11 @@ func NewTestCaCert(version Version, curve Curve, before, after time.Time, networ
// Expiry times are defaulted if you do not pass them in // Expiry times are defaulted if you do not pass them in
func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) { func NewTestCert(v Version, curve Curve, ca Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (Certificate, []byte, []byte, []byte) {
if before.IsZero() { if before.IsZero() {
before = testCertNow.Add(time.Second * -60) before = time.Now().Add(time.Second * -60).Round(time.Second)
} }
if after.IsZero() { if after.IsZero() {
after = testCertNow.Add(time.Second * 60) after = time.Now().Add(time.Second * 60).Round(time.Second)
} }
if len(networks) == 0 { if len(networks) == 0 {
+2 -32
View File
@@ -148,9 +148,6 @@ 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) { func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
k, r := pem.Decode(b) k, r := pem.Decode(b)
if k == nil { if k == nil {
@@ -159,10 +156,10 @@ func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
var expectedLen int var expectedLen int
var curve Curve var curve Curve
switch k.Type { switch k.Type {
case X25519PublicKeyBanner: case X25519PublicKeyBanner, Ed25519PublicKeyBanner:
expectedLen = 32 expectedLen = 32
curve = Curve_CURVE25519 curve = Curve_CURVE25519
case P256PublicKeyBanner: case P256PublicKeyBanner, ECDSAP256PublicKeyBanner:
// Uncompressed // Uncompressed
expectedLen = 65 expectedLen = 65
curve = Curve_P256 curve = Curve_P256
@@ -175,33 +172,6 @@ func UnmarshalPublicKeyFromPEM(b []byte) ([]byte, []byte, Curve, error) {
return k.Bytes, r, curve, nil 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 { func MarshalPrivateKeyToPEM(curve Curve, b []byte) []byte {
switch curve { switch curve {
case Curve_CURVE25519: case Curve_CURVE25519:
+67 -87
View File
@@ -255,6 +255,60 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
func TestUnmarshalPublicKeyFromPEM(t *testing.T) { func TestUnmarshalPublicKeyFromPEM(t *testing.T) {
t.Parallel() t.Parallel()
pubKey := []byte(`# A good key 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----- -----BEGIN NEBULA X25519 PUBLIC KEY-----
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
-----END NEBULA X25519 PUBLIC KEY----- -----END NEBULA X25519 PUBLIC KEY-----
@@ -265,7 +319,7 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
AAAAAAAAAAAAAAAAAAAAAAA= AAAAAAAAAAAAAAAAAAAAAAA=
-----END NEBULA P256 PUBLIC KEY----- -----END NEBULA P256 PUBLIC KEY-----
`) `)
signingKey := []byte(`# A signing key has the wrong scope for this function oldPubP256Key := []byte(`# A good key
-----BEGIN NEBULA ECDSA P256 PUBLIC KEY----- -----BEGIN NEBULA ECDSA P256 PUBLIC KEY-----
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA
AAAAAAAAAAAAAAAAAAAAAAA= AAAAAAAAAAAAAAAAAAAAAAA=
@@ -286,118 +340,44 @@ AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA= AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=
-END NEBULA X25519 PUBLIC KEY-----`) -END NEBULA X25519 PUBLIC KEY-----`)
keyBundle := appendByteSlices(pubKey, pubP256Key, signingKey, shortKey, invalidBanner, invalidPem) keyBundle := appendByteSlices(pubKey, pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem)
// X25519 key // Success test case
k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle) k, rest, curve, err := UnmarshalPublicKeyFromPEM(keyBundle)
assert.Len(t, k, 32) assert.Len(t, k, 32)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, rest, appendByteSlices(pubP256Key, signingKey, shortKey, invalidBanner, invalidPem)) assert.Equal(t, rest, appendByteSlices(pubP256Key, oldPubP256Key, shortKey, invalidBanner, invalidPem))
assert.Equal(t, Curve_CURVE25519, curve) assert.Equal(t, Curve_CURVE25519, curve)
// P256 key // Success test case
k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest) k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
assert.Len(t, k, 65) assert.Len(t, k, 65)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, rest, appendByteSlices(signingKey, shortKey, invalidBanner, invalidPem)) assert.Equal(t, rest, appendByteSlices(oldPubP256Key, shortKey, invalidBanner, invalidPem))
assert.Equal(t, Curve_P256, curve) assert.Equal(t, Curve_P256, curve)
// Reject a signing public key (Ed25519/ECDSA banner) // Success test case
k, rest, _, err = UnmarshalPublicKeyFromPEM(rest) k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
assert.Nil(t, k)
assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
require.EqualError(t, err, "bytes did not contain a proper public key banner")
// Fail due to short key
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, _, 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, _, 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) assert.Len(t, k, 65)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, rest, appendByteSlices(ecdhKey, shortKey, invalidBanner, invalidPem)) assert.Equal(t, rest, appendByteSlices(shortKey, invalidBanner, invalidPem))
assert.Equal(t, Curve_P256, curve) 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 // Fail due to short key
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
assert.Nil(t, k) assert.Nil(t, k)
assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem)) assert.Equal(t, rest, appendByteSlices(invalidBanner, invalidPem))
require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key") require.EqualError(t, err, "key was not 32 bytes, is invalid CURVE25519 public key")
// Fail due to invalid banner // Fail due to invalid banner
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
assert.Nil(t, k) assert.Nil(t, k)
require.EqualError(t, err, "bytes did not contain a proper Ed25519/ECDSA public key banner") require.EqualError(t, err, "bytes did not contain a proper public key banner")
assert.Equal(t, rest, invalidPem) assert.Equal(t, rest, invalidPem)
// Fail due to invalid PEM format, because // Fail due to invalid PEM format, because
// it's missing the requisite pre-encapsulation boundary. // it's missing the requisite pre-encapsulation boundary.
k, rest, _, err = UnmarshalSigningPublicKeyFromPEM(rest) k, rest, curve, err = UnmarshalPublicKeyFromPEM(rest)
assert.Nil(t, k) assert.Nil(t, k)
assert.Equal(t, rest, invalidPem) assert.Equal(t, rest, invalidPem)
require.EqualError(t, err, "input did not contain a valid PEM encoded block") require.EqualError(t, err, "input did not contain a valid PEM encoded block")
+4 -10
View File
@@ -14,12 +14,6 @@ import (
"golang.org/x/crypto/ed25519" "golang.org/x/crypto/ed25519"
) )
// testCertNow is the reference "now" used to derive default before/after times
// in NewTestCaCert and NewTestCert. Holding it fixed for the lifetime of the
// test binary keeps CA and leaf defaults aligned at the same second, so a leaf
// signed with default times can never expire after its CA on a rounding race.
var testCertNow = time.Now().Round(time.Second)
// NewTestCaCert will create a new ca certificate // NewTestCaCert will create a new ca certificate
func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) { func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
var err error var err error
@@ -41,10 +35,10 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
} }
if before.IsZero() { if before.IsZero() {
before = testCertNow.Add(time.Second * -60) before = time.Now().Add(time.Second * -60).Round(time.Second)
} }
if after.IsZero() { if after.IsZero() {
after = testCertNow.Add(time.Second * 60) after = time.Now().Add(time.Second * 60).Round(time.Second)
} }
t := &cert.TBSCertificate{ t := &cert.TBSCertificate{
@@ -77,11 +71,11 @@ func NewTestCaCert(version cert.Version, curve cert.Curve, before, after time.Ti
// Expiry times are defaulted if you do not pass them in // Expiry times are defaulted if you do not pass them in
func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) { func NewTestCert(v cert.Version, curve cert.Curve, ca cert.Certificate, key []byte, name string, before, after time.Time, networks, unsafeNetworks []netip.Prefix, groups []string) (cert.Certificate, []byte, []byte, []byte) {
if before.IsZero() { if before.IsZero() {
before = testCertNow.Add(time.Second * -60) before = time.Now().Add(time.Second * -60).Round(time.Second)
} }
if after.IsZero() { if after.IsZero() {
after = testCertNow.Add(time.Second * 60) after = time.Now().Add(time.Second * 60).Round(time.Second)
} }
var pub, priv []byte var pub, priv []byte
+5 -30
View File
@@ -97,19 +97,6 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
if err = mustFlagString("out-key", cf.outKeyPath); err != nil { if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
return err return err
} }
} else {
// out-key is meaningless under PKCS#11 because the private key never
// leaves the HSM; reject it so we never silently accept or claim a
// stdout slot for it.
outKeySet := false
cf.set.Visit(func(f *flag.Flag) {
if f.Name == "out-key" {
outKeySet = true
}
})
if outKeySet {
return newHelpErrorf("cannot set -out-key with -pkcs11")
}
} }
if err := mustFlagString("out-crt", cf.outCertPath); err != nil { if err := mustFlagString("out-crt", cf.outCertPath); err != nil {
return err return err
@@ -184,21 +171,12 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
} }
} }
var claims ioClaims
if err := reserveOutputs(&claims,
"out-key", *cf.outKeyPath,
"out-crt", *cf.outCertPath,
"out-qr", *cf.outQRPath,
); err != nil {
return err
}
var passphrase []byte var passphrase []byte
if !isP11 && *cf.encryption { if !isP11 && *cf.encryption {
passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE")) passphrase = []byte(os.Getenv("NEBULA_CA_PASSPHRASE"))
if len(passphrase) == 0 { if len(passphrase) == 0 {
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
errOut.Write([]byte("Enter passphrase: ")) out.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword() passphrase, err = pr.ReadPassword()
if err == ErrNoTerminal { if err == ErrNoTerminal {
@@ -283,17 +261,15 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
Curve: curve, Curve: curve,
} }
if !isP11 && !isStdio(*cf.outKeyPath) { if !isP11 {
if _, err := os.Stat(*cf.outKeyPath); err == nil { if _, err := os.Stat(*cf.outKeyPath); err == nil {
return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath) return fmt.Errorf("refusing to overwrite existing CA key: %s", *cf.outKeyPath)
} }
} }
if !isStdio(*cf.outCertPath) {
if _, err := os.Stat(*cf.outCertPath); err == nil { if _, err := os.Stat(*cf.outCertPath); err == nil {
return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath) return fmt.Errorf("refusing to overwrite existing CA cert: %s", *cf.outCertPath)
} }
}
var c cert.Certificate var c cert.Certificate
var b []byte var b []byte
@@ -318,7 +294,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv) b = cert.MarshalSigningPrivateKeyToPEM(curve, rawPriv)
} }
err = writeOutput(*cf.outKeyPath, b, 0600, out) err = os.WriteFile(*cf.outKeyPath, b, 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-key: %s", err) return fmt.Errorf("error while writing out-key: %s", err)
} }
@@ -329,7 +305,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
return fmt.Errorf("error while marshalling certificate: %s", err) return fmt.Errorf("error while marshalling certificate: %s", err)
} }
err = writeOutput(*cf.outCertPath, b, 0600, out) err = os.WriteFile(*cf.outCertPath, b, 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-crt: %s", err) return fmt.Errorf("error while writing out-crt: %s", err)
} }
@@ -340,7 +316,7 @@ func ca(args []string, out io.Writer, errOut io.Writer, pr PasswordReader) error
return fmt.Errorf("error while generating qr code: %s", err) return fmt.Errorf("error while generating qr code: %s", err)
} }
err = writeOutput(*cf.outQRPath, b, 0600, out) err = os.WriteFile(*cf.outQRPath, b, 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-qr: %s", err) return fmt.Errorf("error while writing out-qr: %s", err)
} }
@@ -356,7 +332,6 @@ func caSummary() string {
func caHelp(out io.Writer) { func caHelp(out io.Writer) {
cf := newCaFlags() cf := newCaFlags()
out.Write([]byte("Usage of " + os.Args[0] + " " + caSummary() + "\n")) out.Write([]byte("Usage of " + os.Args[0] + " " + caSummary() + "\n"))
out.Write([]byte(stdioHelpText))
cf.set.SetOutput(out) cf.set.SetOutput(out)
cf.set.PrintDefaults() cf.set.PrintDefaults()
} }
+7 -72
View File
@@ -27,7 +27,6 @@ func Test_caHelp(t *testing.T) {
assert.Equal( assert.Equal(
t, t,
"Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+ "Usage of "+os.Args[0]+" ca <flags>: create a self signed certificate authority\n"+
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
" -argon-iterations uint\n"+ " -argon-iterations uint\n"+
" \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+ " \tOptional: Argon2 iterations parameter used for encrypted private key passphrase (default 1)\n"+
" -argon-memory uint\n"+ " -argon-memory uint\n"+
@@ -85,7 +84,7 @@ func Test_ca(t *testing.T) {
err: nil, err: nil,
} }
pwPromptEB := "Enter passphrase: " pwPromptOb := "Enter passphrase: "
// required args // required args
assertHelpError(t, ca( assertHelpError(t, ca(
@@ -169,8 +168,8 @@ func Test_ca(t *testing.T) {
eb.Reset() eb.Reset()
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()} args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
require.NoError(t, ca(args, ob, eb, testpw)) require.NoError(t, ca(args, ob, eb, testpw))
assert.Empty(t, ob.String()) assert.Equal(t, pwPromptOb, ob.String())
assert.Equal(t, pwPromptEB, eb.String()) assert.Empty(t, eb.String())
// test encrypted key with passphrase environment variable // test encrypted key with passphrase environment variable
os.Remove(keyF.Name()) os.Remove(keyF.Name())
@@ -208,8 +207,8 @@ func Test_ca(t *testing.T) {
eb.Reset() eb.Reset()
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()} args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
require.Error(t, ca(args, ob, eb, errpw)) require.Error(t, ca(args, ob, eb, errpw))
assert.Empty(t, ob.String()) assert.Equal(t, pwPromptOb, ob.String())
assert.Equal(t, pwPromptEB, eb.String()) assert.Empty(t, eb.String())
// test when user fails to enter a password // test when user fails to enter a password
os.Remove(keyF.Name()) os.Remove(keyF.Name())
@@ -218,8 +217,8 @@ func Test_ca(t *testing.T) {
eb.Reset() eb.Reset()
args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()} args = []string{"-version", "1", "-encrypt", "-name", "test", "-duration", "100m", "-groups", "1,2,3,4,5", "-out-crt", crtF.Name(), "-out-key", keyF.Name()}
require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext") require.EqualError(t, ca(args, ob, eb, nopw), "no passphrase specified, remove -encrypt flag to write out-key in plaintext")
assert.Empty(t, ob.String()) assert.Equal(t, strings.Repeat(pwPromptOb, 5), ob.String()) // prompts 5 times before giving up
assert.Equal(t, strings.Repeat(pwPromptEB, 5), eb.String()) // prompts 5 times before giving up assert.Empty(t, eb.String())
// create valid cert/key for overwrite tests // create valid cert/key for overwrite tests
os.Remove(keyF.Name()) os.Remove(keyF.Name())
@@ -248,67 +247,3 @@ func Test_ca(t *testing.T) {
os.Remove(keyF.Name()) os.Remove(keyF.Name())
} }
func Test_ca_stdio(t *testing.T) {
nopw := &StubPasswordReader{}
keyF, err := os.CreateTemp("", "ca.key")
require.NoError(t, err)
os.Remove(keyF.Name())
defer os.Remove(keyF.Name())
crtF, err := os.CreateTemp("", "ca.crt")
require.NoError(t, err)
os.Remove(crtF.Name())
defer os.Remove(crtF.Name())
// out-crt on stdout, out-key on disk
ob := &bytes.Buffer{}
eb := &bytes.Buffer{}
require.NoError(t, ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", "-", "-out-key", keyF.Name()}, ob, eb, nopw))
assert.Empty(t, eb.String())
c, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
require.NoError(t, err)
assert.True(t, c.IsCA())
assert.Equal(t, "test-ca", c.Name())
// out-key on stdout, out-crt on disk
os.Remove(keyF.Name())
ob.Reset()
eb.Reset()
require.NoError(t, ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", crtF.Name(), "-out-key", "-"}, ob, eb, nopw))
assert.Empty(t, eb.String())
_, _, curve, err := cert.UnmarshalSigningPrivateKeyFromPEM(ob.Bytes())
require.NoError(t, err)
assert.Equal(t, cert.Curve_CURVE25519, curve)
// dual stdout is rejected up front
os.Remove(crtF.Name())
ob.Reset()
eb.Reset()
require.EqualError(t,
ca([]string{"-name", "test-ca", "-duration", "1h", "-out-crt", "-", "-out-key", "-"}, ob, eb, nopw),
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
assert.Empty(t, ob.String())
// an output conflict combined with -encrypt must error BEFORE prompting
// for a passphrase; pr would record any read attempt
tracker := &trackingPasswordReader{}
ob.Reset()
eb.Reset()
require.EqualError(t,
ca([]string{"-name", "test-ca", "-duration", "1h", "-encrypt", "-out-crt", "-", "-out-key", "-"}, ob, eb, tracker),
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
assert.Empty(t, ob.String())
assert.Empty(t, eb.String())
assert.Zero(t, tracker.calls, "passphrase prompt should not have been called")
}
type trackingPasswordReader struct {
calls int
}
func (pr *trackingPasswordReader) ReadPassword() ([]byte, error) {
pr.calls++
return []byte(""), nil
}
+2 -13
View File
@@ -42,8 +42,6 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
if err = mustFlagString("out-key", cf.outKeyPath); err != nil { if err = mustFlagString("out-key", cf.outKeyPath); err != nil {
return err return err
} }
} else if *cf.outKeyPath != "" {
return newHelpErrorf("cannot set -out-key with -pkcs11")
} }
if err = mustFlagString("out-pub", cf.outPubPath); err != nil { if err = mustFlagString("out-pub", cf.outPubPath); err != nil {
return err return err
@@ -71,14 +69,6 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
} }
} }
var claims ioClaims
if err := reserveOutputs(&claims,
"out-key", *cf.outKeyPath,
"out-pub", *cf.outPubPath,
); err != nil {
return err
}
if isP11 { if isP11 {
p11Client, err := pkclient.FromUrl(*cf.p11url) p11Client, err := pkclient.FromUrl(*cf.p11url)
if err != nil { if err != nil {
@@ -92,12 +82,12 @@ func keygen(args []string, out io.Writer, errOut io.Writer) error {
return fmt.Errorf("error while getting public key: %w", err) return fmt.Errorf("error while getting public key: %w", err)
} }
} else { } else {
err = writeOutput(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out) err = os.WriteFile(*cf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-key: %s", err) return fmt.Errorf("error while writing out-key: %s", err)
} }
} }
err = writeOutput(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600, out) err = os.WriteFile(*cf.outPubPath, cert.MarshalPublicKeyToPEM(curve, pub), 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-pub: %s", err) return fmt.Errorf("error while writing out-pub: %s", err)
} }
@@ -112,7 +102,6 @@ func keygenSummary() string {
func keygenHelp(out io.Writer) { func keygenHelp(out io.Writer) {
cf := newKeygenFlags() cf := newKeygenFlags()
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n")) _, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + keygenSummary() + "\n"))
_, _ = out.Write([]byte(stdioHelpText))
cf.set.SetOutput(out) cf.set.SetOutput(out)
cf.set.PrintDefaults() cf.set.PrintDefaults()
} }
-41
View File
@@ -20,7 +20,6 @@ func Test_keygenHelp(t *testing.T) {
assert.Equal( assert.Equal(
t, t,
"Usage of "+os.Args[0]+" keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`\n"+ "Usage of "+os.Args[0]+" keygen <flags>: create a public/private key pair. the public key can be passed to `nebula-cert sign`\n"+
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
" -curve string\n"+ " -curve string\n"+
" \tECDH Curve (25519, P256) (default \"25519\")\n"+ " \tECDH Curve (25519, P256) (default \"25519\")\n"+
" -out-key string\n"+ " -out-key string\n"+
@@ -94,43 +93,3 @@ func Test_keygen(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
assert.Len(t, lPub, 32) assert.Len(t, lPub, 32)
} }
func Test_keygen_stdio(t *testing.T) {
keyF, err := os.CreateTemp("", "test.key")
require.NoError(t, err)
os.Remove(keyF.Name())
defer os.Remove(keyF.Name())
pubF, err := os.CreateTemp("", "test.pub")
require.NoError(t, err)
os.Remove(pubF.Name())
defer os.Remove(pubF.Name())
// out-pub on stdout, out-key on disk
ob := &bytes.Buffer{}
eb := &bytes.Buffer{}
require.NoError(t, keygen([]string{"-out-pub", "-", "-out-key", keyF.Name()}, ob, eb))
assert.Empty(t, eb.String())
lPub, _, curve, err := cert.UnmarshalPublicKeyFromPEM(ob.Bytes())
require.NoError(t, err)
assert.Equal(t, cert.Curve_CURVE25519, curve)
assert.Len(t, lPub, 32)
// out-key on stdout, out-pub on disk
os.Remove(keyF.Name())
ob.Reset()
eb.Reset()
require.NoError(t, keygen([]string{"-out-pub", pubF.Name(), "-out-key", "-"}, ob, eb))
assert.Empty(t, eb.String())
lKey, _, curve, err := cert.UnmarshalPrivateKeyFromPEM(ob.Bytes())
require.NoError(t, err)
assert.Equal(t, cert.Curve_CURVE25519, curve)
assert.Len(t, lKey, 32)
// both on stdout is a conflict caught up front
ob.Reset()
eb.Reset()
require.EqualError(t, keygen([]string{"-out-pub", "-", "-out-key", "-"}, ob, eb),
`-out-key and -out-pub both set to "-", only one output may write to stdout`)
assert.Empty(t, ob.String())
}
+1 -3
View File
@@ -22,9 +22,7 @@ func (pr StdinPasswordReader) ReadPassword() ([]byte, error) {
} }
password, err := term.ReadPassword(int(os.Stdin.Fd())) password, err := term.ReadPassword(int(os.Stdin.Fd()))
// Terminal echo is off while reading, so the user's Enter key does not fmt.Println()
// produce a visible newline. Emit one on stderr to match the prompt.
fmt.Fprintln(os.Stderr)
return password, err return password, err
} }
+3 -18
View File
@@ -40,23 +40,11 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
return err return err
} }
var claims ioClaims rawCert, err := os.ReadFile(*pf.path)
if err := reserveInputs(&claims, "path", *pf.path); err != nil {
return err
}
if err := reserveOutputs(&claims, "out-qr", *pf.outQRPath); err != nil {
return err
}
rawCert, err := readInput("path", *pf.path, &claims)
if err != nil { if err != nil {
return fmt.Errorf("unable to read cert; %s", err) return fmt.Errorf("unable to read cert; %s", err)
} }
// When the QR is going to stdout, suppress the human-readable text/json
// output so the binary stream is not contaminated.
qrToStdout := isStdio(*pf.outQRPath)
var c cert.Certificate var c cert.Certificate
var qrBytes []byte var qrBytes []byte
part := 0 part := 0
@@ -69,14 +57,12 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
return fmt.Errorf("error while unmarshaling cert: %s", err) return fmt.Errorf("error while unmarshaling cert: %s", err)
} }
if !qrToStdout {
if *pf.json { if *pf.json {
jsonCerts = append(jsonCerts, c) jsonCerts = append(jsonCerts, c)
} else { } else {
_, _ = out.Write([]byte(c.String())) _, _ = out.Write([]byte(c.String()))
_, _ = out.Write([]byte("\n")) _, _ = out.Write([]byte("\n"))
} }
}
if *pf.outQRPath != "" { if *pf.outQRPath != "" {
b, err := c.MarshalPEM() b, err := c.MarshalPEM()
@@ -93,7 +79,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
part++ part++
} }
if *pf.json && !qrToStdout { if *pf.json {
b, _ := json.Marshal(jsonCerts) b, _ := json.Marshal(jsonCerts)
_, _ = out.Write(b) _, _ = out.Write(b)
_, _ = out.Write([]byte("\n")) _, _ = out.Write([]byte("\n"))
@@ -105,7 +91,7 @@ func printCert(args []string, out io.Writer, errOut io.Writer) error {
return fmt.Errorf("error while generating qr code: %s", err) return fmt.Errorf("error while generating qr code: %s", err)
} }
err = writeOutput(*pf.outQRPath, b, 0600, out) err = os.WriteFile(*pf.outQRPath, b, 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-qr: %s", err) return fmt.Errorf("error while writing out-qr: %s", err)
} }
@@ -121,7 +107,6 @@ func printSummary() string {
func printHelp(out io.Writer) { func printHelp(out io.Writer) {
pf := newPrintFlags() pf := newPrintFlags()
out.Write([]byte("Usage of " + os.Args[0] + " " + printSummary() + "\n")) out.Write([]byte("Usage of " + os.Args[0] + " " + printSummary() + "\n"))
out.Write([]byte(stdioHelpText))
pf.set.SetOutput(out) pf.set.SetOutput(out)
pf.set.PrintDefaults() pf.set.PrintDefaults()
} }
-39
View File
@@ -25,7 +25,6 @@ func Test_printHelp(t *testing.T) {
assert.Equal( assert.Equal(
t, t,
"Usage of "+os.Args[0]+" print <flags>: prints details about a certificate\n"+ "Usage of "+os.Args[0]+" print <flags>: prints details about a certificate\n"+
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
" -json\n"+ " -json\n"+
" \tOptional: outputs certificates in json format\n"+ " \tOptional: outputs certificates in json format\n"+
" -out-qr string\n"+ " -out-qr string\n"+
@@ -179,44 +178,6 @@ func Test_printCert(t *testing.T) {
ob.String(), ob.String(),
) )
assert.Empty(t, eb.String()) assert.Empty(t, eb.String())
// read cert from stdin
ob.Reset()
eb.Reset()
withStdin(t, bytes.NewReader(p))
err = printCert([]string{"-json", "-path", "-"}, ob, eb)
require.NoError(t, err)
assert.Equal(
t,
`[{"details":{"curve":"CURVE25519","groups":["hi"],"isCa":false,"issuer":"`+c.Issuer()+`","name":"test","networks":["10.0.0.123/8"],"notAfter":"0001-01-01T00:00:00Z","notBefore":"0001-01-01T00:00:00Z","publicKey":"`+pk+`","unsafeNetworks":[]},"fingerprint":"`+fp+`","signature":"`+sig+`","version":1}]
`,
ob.String(),
)
assert.Empty(t, eb.String())
// -out-qr - sends only the PNG to stdout, suppressing the cert dump
ob.Reset()
eb.Reset()
withStdin(t, bytes.NewReader(p))
err = printCert([]string{"-path", "-", "-out-qr", "-"}, ob, eb)
require.NoError(t, err)
assert.Empty(t, eb.String())
stdout := ob.Bytes()
require.NotEmpty(t, stdout)
// PNG magic, no PEM/JSON noise prepended
assert.Equal(t, []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a}, stdout[:8])
assert.NotContains(t, string(stdout), "NebulaCertificate")
assert.NotContains(t, string(stdout), `"details"`)
// json + out-qr - still suppresses json
ob.Reset()
eb.Reset()
withStdin(t, bytes.NewReader(p))
err = printCert([]string{"-json", "-path", "-", "-out-qr", "-"}, ob, eb)
require.NoError(t, err)
assert.Empty(t, eb.String())
assert.Equal(t, []byte{0x89, 'P', 'N', 'G'}, ob.Bytes()[:4])
assert.NotContains(t, ob.String(), `"details"`)
} }
// NewTestCaCert will generate a CA cert // NewTestCaCert will generate a CA cert
+16 -38
View File
@@ -85,9 +85,6 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" { if !isP11 && *sf.inPubPath != "" && *sf.outKeyPath != "" {
return newHelpErrorf("cannot set both -in-pub and -out-key") return newHelpErrorf("cannot set both -in-pub and -out-key")
} }
if isP11 && *sf.outKeyPath != "" {
return newHelpErrorf("cannot set -out-key with -pkcs11")
}
var v4Networks []netip.Prefix var v4Networks []netip.Prefix
var v6Networks []netip.Prefix var v6Networks []netip.Prefix
@@ -105,35 +102,13 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2) return newHelpErrorf("-version must be either %v or %v", cert.Version1, cert.Version2)
} }
if *sf.outKeyPath == "" {
*sf.outKeyPath = *sf.name + ".key"
}
if *sf.outCertPath == "" {
*sf.outCertPath = *sf.name + ".crt"
}
var claims ioClaims
if err := reserveInputs(&claims,
"ca-key", *sf.caKeyPath,
"ca-crt", *sf.caCertPath,
"in-pub", *sf.inPubPath,
); err != nil {
return err
}
if err := reserveOutputs(&claims,
"out-key", *sf.outKeyPath,
"out-crt", *sf.outCertPath,
"out-qr", *sf.outQRPath,
); err != nil {
return err
}
var curve cert.Curve var curve cert.Curve
var caKey []byte var caKey []byte
if !isP11 { if !isP11 {
var rawCAKey []byte var rawCAKey []byte
rawCAKey, err = readInput("ca-key", *sf.caKeyPath, &claims) rawCAKey, err := os.ReadFile(*sf.caKeyPath)
if err != nil { if err != nil {
return fmt.Errorf("error while reading ca-key: %s", err) return fmt.Errorf("error while reading ca-key: %s", err)
} }
@@ -146,7 +121,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
if len(passphrase) == 0 { if len(passphrase) == 0 {
// ask for a passphrase until we get one // ask for a passphrase until we get one
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
errOut.Write([]byte("Enter passphrase: ")) out.Write([]byte("Enter passphrase: "))
passphrase, err = pr.ReadPassword() passphrase, err = pr.ReadPassword()
if errors.Is(err, ErrNoTerminal) { if errors.Is(err, ErrNoTerminal) {
@@ -172,7 +147,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
} }
} }
rawCACert, err := readInput("ca-crt", *sf.caCertPath, &claims) rawCACert, err := os.ReadFile(*sf.caCertPath)
if err != nil { if err != nil {
return fmt.Errorf("error while reading ca-crt: %s", err) return fmt.Errorf("error while reading ca-crt: %s", err)
} }
@@ -270,7 +245,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
if *sf.inPubPath != "" { if *sf.inPubPath != "" {
var pubCurve cert.Curve var pubCurve cert.Curve
rawPub, err := readInput("in-pub", *sf.inPubPath, &claims) rawPub, err := os.ReadFile(*sf.inPubPath)
if err != nil { if err != nil {
return fmt.Errorf("error while reading in-pub: %s", err) return fmt.Errorf("error while reading in-pub: %s", err)
} }
@@ -291,11 +266,17 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
pub, rawPriv = newKeypair(curve) pub, rawPriv = newKeypair(curve)
} }
if !isStdio(*sf.outCertPath) { if *sf.outKeyPath == "" {
*sf.outKeyPath = *sf.name + ".key"
}
if *sf.outCertPath == "" {
*sf.outCertPath = *sf.name + ".crt"
}
if _, err := os.Stat(*sf.outCertPath); err == nil { if _, err := os.Stat(*sf.outCertPath); err == nil {
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath) return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
} }
}
var crts []cert.Certificate var crts []cert.Certificate
@@ -379,13 +360,11 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
} }
if !isP11 && *sf.inPubPath == "" { if !isP11 && *sf.inPubPath == "" {
if !isStdio(*sf.outKeyPath) {
if _, err := os.Stat(*sf.outKeyPath); err == nil { if _, err := os.Stat(*sf.outKeyPath); err == nil {
return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath) return fmt.Errorf("refusing to overwrite existing key: %s", *sf.outKeyPath)
} }
}
err = writeOutput(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600, out) err = os.WriteFile(*sf.outKeyPath, cert.MarshalPrivateKeyToPEM(curve, rawPriv), 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-key: %s", err) return fmt.Errorf("error while writing out-key: %s", err)
} }
@@ -400,7 +379,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
b = append(b, sb...) b = append(b, sb...)
} }
err = writeOutput(*sf.outCertPath, b, 0600, out) err = os.WriteFile(*sf.outCertPath, b, 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-crt: %s", err) return fmt.Errorf("error while writing out-crt: %s", err)
} }
@@ -411,7 +390,7 @@ func signCert(args []string, out io.Writer, errOut io.Writer, pr PasswordReader)
return fmt.Errorf("error while generating qr code: %s", err) return fmt.Errorf("error while generating qr code: %s", err)
} }
err = writeOutput(*sf.outQRPath, b, 0600, out) err = os.WriteFile(*sf.outQRPath, b, 0600)
if err != nil { if err != nil {
return fmt.Errorf("error while writing out-qr: %s", err) return fmt.Errorf("error while writing out-qr: %s", err)
} }
@@ -461,7 +440,6 @@ func signSummary() string {
func signHelp(out io.Writer) { func signHelp(out io.Writer) {
sf := newSignFlags() sf := newSignFlags()
out.Write([]byte("Usage of " + os.Args[0] + " " + signSummary() + "\n")) out.Write([]byte("Usage of " + os.Args[0] + " " + signSummary() + "\n"))
out.Write([]byte(stdioHelpText))
sf.set.SetOutput(out) sf.set.SetOutput(out)
sf.set.PrintDefaults() sf.set.PrintDefaults()
} }
+7 -111
View File
@@ -27,7 +27,6 @@ func Test_signHelp(t *testing.T) {
assert.Equal( assert.Equal(
t, t,
"Usage of "+os.Args[0]+" sign <flags>: create and sign a certificate\n"+ "Usage of "+os.Args[0]+" sign <flags>: create and sign a certificate\n"+
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
" -ca-crt string\n"+ " -ca-crt string\n"+
" \tOptional: path to the signing CA cert (default \"ca.crt\")\n"+ " \tOptional: path to the signing CA cert (default \"ca.crt\")\n"+
" -ca-key string\n"+ " -ca-key string\n"+
@@ -377,18 +376,15 @@ func Test_signCert(t *testing.T) {
// test with the proper password // test with the proper password
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"} args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
require.NoError(t, signCert(args, ob, eb, testpw)) require.NoError(t, signCert(args, ob, eb, testpw))
assert.Empty(t, ob.String()) assert.Equal(t, "Enter passphrase: ", ob.String())
assert.Equal(t, "Enter passphrase: ", eb.String()) assert.Empty(t, eb.String())
// test with the proper password in the environment // test with the proper password in the environment
os.Remove(crtF.Name()) os.Remove(crtF.Name())
os.Remove(keyF.Name()) os.Remove(keyF.Name())
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"} args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase)) os.Setenv("NEBULA_CA_PASSPHRASE", string(passphrase))
ob.Reset()
eb.Reset()
require.NoError(t, signCert(args, ob, eb, testpw)) require.NoError(t, signCert(args, ob, eb, testpw))
assert.Empty(t, ob.String())
assert.Empty(t, eb.String()) assert.Empty(t, eb.String())
os.Setenv("NEBULA_CA_PASSPHRASE", "") os.Setenv("NEBULA_CA_PASSPHRASE", "")
@@ -399,8 +395,8 @@ func Test_signCert(t *testing.T) {
testpw.password = []byte("invalid password") testpw.password = []byte("invalid password")
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"} args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
require.Error(t, signCert(args, ob, eb, testpw)) require.Error(t, signCert(args, ob, eb, testpw))
assert.Empty(t, ob.String()) assert.Equal(t, "Enter passphrase: ", ob.String())
assert.Equal(t, "Enter passphrase: ", eb.String()) assert.Empty(t, eb.String())
// test with the wrong password in environment // test with the wrong password in environment
ob.Reset() ob.Reset()
@@ -420,8 +416,8 @@ func Test_signCert(t *testing.T) {
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"} args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
require.Error(t, signCert(args, ob, eb, nopw)) require.Error(t, signCert(args, ob, eb, nopw))
// normally the user hitting enter on the prompt would add newlines between these // normally the user hitting enter on the prompt would add newlines between these
assert.Empty(t, ob.String()) assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", ob.String())
assert.Equal(t, "Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: Enter passphrase: ", eb.String()) assert.Empty(t, eb.String())
// test an error condition // test an error condition
ob.Reset() ob.Reset()
@@ -429,106 +425,6 @@ func Test_signCert(t *testing.T) {
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"} args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", keyF.Name(), "-duration", "100m", "-subnets", "10.1.1.1/32, , 10.2.2.2/32 , , ,, 10.5.5.5/32", "-groups", "1,, 2 , ,,,3,4,5"}
require.Error(t, signCert(args, ob, eb, errpw)) require.Error(t, signCert(args, ob, eb, errpw))
assert.Empty(t, ob.String()) assert.Equal(t, "Enter passphrase: ", ob.String())
assert.Equal(t, "Enter passphrase: ", eb.String())
}
func Test_signCert_stdio(t *testing.T) {
nopw := &StubPasswordReader{
password: []byte(""),
err: nil,
}
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
rawCAKey := cert.MarshalSigningPrivateKeyToPEM(cert.Curve_CURVE25519, caPriv)
ca, _ := NewTestCaCert("ca", caPub, caPriv, time.Now(), time.Now().Add(time.Minute*200), nil, nil, nil)
rawCACrt, _ := ca.MarshalPEM()
caCrtF, err := os.CreateTemp("", "sign-cert.crt")
require.NoError(t, err)
defer os.Remove(caCrtF.Name())
caCrtF.Write(rawCACrt)
caKeyF, err := os.CreateTemp("", "sign-cert.key")
require.NoError(t, err)
defer os.Remove(caKeyF.Name())
caKeyF.Write(rawCAKey)
keyF, err := os.CreateTemp("", "sign.key")
require.NoError(t, err)
os.Remove(keyF.Name())
defer os.Remove(keyF.Name())
// ca-key on stdin, cert to stdout
withStdin(t, bytes.NewReader(rawCAKey))
ob := &bytes.Buffer{}
eb := &bytes.Buffer{}
args := []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "-", "-out-key", keyF.Name(), "-duration", "100m"}
require.NoError(t, signCert(args, ob, eb, nopw))
assert.Empty(t, eb.String()) assert.Empty(t, eb.String())
lCrt, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
require.NoError(t, err)
assert.Equal(t, "stdin-test", lCrt.Name())
assert.True(t, lCrt.CheckSignature(caPub))
// two flags reading from stdin should error before any read attempt;
// otherwise an interactive shell would hang on io.ReadAll
stdinIn := bytes.NewReader(rawCAKey)
withStdin(t, stdinIn)
ob.Reset()
eb.Reset()
args = []string{"-version", "1", "-ca-crt", "-", "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "nope", "-out-key", "nope", "-duration", "100m"}
require.EqualError(t, signCert(args, ob, eb, nopw),
`-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
assert.Equal(t, len(rawCAKey), stdinIn.Len(), "stdin should be untouched when conflict is caught up front")
// two flags writing to stdout should error before any output is written
// AND before stdin is consumed
stdinR := bytes.NewReader(rawCAKey)
withStdin(t, stdinR)
ob.Reset()
eb.Reset()
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", "-", "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", "-", "-out-key", "-", "-duration", "100m"}
require.EqualError(t, signCert(args, ob, eb, nopw),
`-out-key and -out-crt both set to "-", only one output may write to stdout`)
assert.Empty(t, ob.String())
// stdin should be untouched because the conflict was caught up front
assert.Equal(t, len(rawCAKey), stdinR.Len())
// out-key on stdout, cert on disk
keyF2, err := os.CreateTemp("", "sign.key")
require.NoError(t, err)
os.Remove(keyF2.Name())
defer os.Remove(keyF2.Name())
crtF, err := os.CreateTemp("", "sign.crt")
require.NoError(t, err)
os.Remove(crtF.Name())
defer os.Remove(crtF.Name())
ob.Reset()
eb.Reset()
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "stdin-test", "-ip", "1.1.1.1/24", "-out-crt", crtF.Name(), "-out-key", "-", "-duration", "100m"}
require.NoError(t, signCert(args, ob, eb, nopw))
assert.Empty(t, eb.String())
_, _, curve, err := cert.UnmarshalPrivateKeyFromPEM(ob.Bytes())
require.NoError(t, err)
assert.Equal(t, cert.Curve_CURVE25519, curve)
// in-pub on stdin (caller already has a keypair, only the cert is generated)
inPub, _ := x25519Keypair()
rawInPub := cert.MarshalPublicKeyToPEM(cert.Curve_CURVE25519, inPub)
withStdin(t, bytes.NewReader(rawInPub))
os.Remove(crtF.Name())
ob.Reset()
eb.Reset()
args = []string{"-version", "1", "-ca-crt", caCrtF.Name(), "-ca-key", caKeyF.Name(), "-name", "in-pub-test", "-ip", "1.1.1.1/24", "-in-pub", "-", "-out-crt", "-", "-duration", "100m"}
require.NoError(t, signCert(args, ob, eb, nopw))
assert.Empty(t, eb.String())
stdinCrt, _, err := cert.UnmarshalCertificateFromPEM(ob.Bytes())
require.NoError(t, err)
assert.Equal(t, "in-pub-test", stdinCrt.Name())
assert.Equal(t, inPub, stdinCrt.PublicKey())
} }
-117
View File
@@ -1,117 +0,0 @@
package main
import (
"fmt"
"io"
"os"
)
// stdioPath is the special path value that selects stdin (for inputs) or
// stdout (for outputs) instead of a file on disk.
const stdioPath = "-"
// stdioHelpText is rendered just under the Usage line of each subcommand
// help so the - convention is documented once instead of on every flag.
const stdioHelpText = " Pass \"-\" to any path flag to read from stdin or write to stdout.\n"
// stdinReader is the source used when an input flag is set to "-".
// It is a package level var so tests can swap in a deterministic reader.
// Tests that mutate stdinReader cannot run with t.Parallel().
var stdinReader io.Reader = os.Stdin
// ioClaims tracks which flags have claimed stdin and stdout during a single
// command invocation so we can refuse a second flag asking for the same
// stream.
type ioClaims struct {
in string
out string
}
func (c *ioClaims) claimIn(flagName string) error {
if c.in != "" && c.in != flagName {
return fmt.Errorf("-%s and -%s both set to %q, only one input may read from stdin", c.in, flagName, stdioPath)
}
c.in = flagName
return nil
}
func (c *ioClaims) claimOut(flagName string) error {
if c.out != "" && c.out != flagName {
return fmt.Errorf("-%s and -%s both set to %q, only one output may write to stdout", c.out, flagName, stdioPath)
}
c.out = flagName
return nil
}
// reserveInputs walks alternating (flagName, path) pairs and claims stdin
// for any path equal to stdioPath. It must be called before any input is
// read so a conflict can be reported immediately instead of blocking on
// io.ReadAll while waiting for input that will never arrive.
func reserveInputs(claims *ioClaims, pairs ...string) error {
return reserveStdio(claims, "reserveInputs", (*ioClaims).claimIn, pairs)
}
// reserveOutputs walks alternating (flagName, path) pairs and claims stdout
// for any path equal to stdioPath. It must be called before any output is
// written so a conflict cannot leave one stream half written before the
// second flag fails.
func reserveOutputs(claims *ioClaims, pairs ...string) error {
return reserveStdio(claims, "reserveOutputs", (*ioClaims).claimOut, pairs)
}
func reserveStdio(claims *ioClaims, who string, claim func(*ioClaims, string) error, pairs []string) error {
if len(pairs)%2 != 0 {
panic(who + " requires alternating name, path pairs")
}
for i := 0; i < len(pairs); i += 2 {
name, path := pairs[i], pairs[i+1]
if path != stdioPath {
continue
}
if err := claim(claims, name); err != nil {
return err
}
}
return nil
}
// readInput returns the bytes referenced by path, reading from stdin when
// path is stdioPath.
func readInput(flagName, path string, claims *ioClaims) ([]byte, error) {
if path == stdioPath {
if err := claims.claimIn(flagName); err != nil {
return nil, err
}
return io.ReadAll(stdinReader)
}
return os.ReadFile(path)
}
// openInput returns a reader for path. When path is stdioPath the returned
// reader wraps stdin and Close is a no-op.
func openInput(flagName, path string, claims *ioClaims) (io.ReadCloser, error) {
if path == stdioPath {
if err := claims.claimIn(flagName); err != nil {
return nil, err
}
return io.NopCloser(stdinReader), nil
}
return os.Open(path)
}
// writeOutput writes data to path, or to stdout when path is stdioPath. perm
// is only used for file output. The caller must have already claimed stdout
// via reserveOutputs before invoking with stdioPath.
func writeOutput(path string, data []byte, perm os.FileMode, stdout io.Writer) error {
if path == stdioPath {
_, err := stdout.Write(data)
return err
}
return os.WriteFile(path, data, perm)
}
// isStdio reports whether path is the stdio sentinel and so should skip
// existence checks like "refuse to overwrite".
func isStdio(path string) bool {
return path == stdioPath
}
-167
View File
@@ -1,167 +0,0 @@
package main
import (
"bytes"
"io"
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// withStdin temporarily replaces stdinReader for the duration of t.
func withStdin(t *testing.T, r io.Reader) {
t.Helper()
prev := stdinReader
stdinReader = r
t.Cleanup(func() { stdinReader = prev })
}
func Test_readInput_stdin(t *testing.T) {
withStdin(t, bytes.NewBufferString("hello"))
var claims ioClaims
got, err := readInput("path", "-", &claims)
require.NoError(t, err)
assert.Equal(t, []byte("hello"), got)
assert.Equal(t, "path", claims.in)
}
func Test_readInput_file(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "f")
require.NoError(t, os.WriteFile(p, []byte("file"), 0600))
var claims ioClaims
got, err := readInput("path", p, &claims)
require.NoError(t, err)
assert.Equal(t, []byte("file"), got)
assert.Empty(t, claims.in)
}
func Test_readInput_doubleStdinErrors(t *testing.T) {
withStdin(t, bytes.NewBufferString("hello"))
var claims ioClaims
_, err := readInput("ca-key", "-", &claims)
require.NoError(t, err)
_, err = readInput("ca-crt", "-", &claims)
require.EqualError(t, err, `-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
}
func Test_openInput_stdin(t *testing.T) {
withStdin(t, bytes.NewBufferString("hi"))
var claims ioClaims
r, err := openInput("ca", "-", &claims)
require.NoError(t, err)
defer r.Close()
b, err := io.ReadAll(r)
require.NoError(t, err)
assert.Equal(t, []byte("hi"), b)
}
func Test_openInput_doubleStdinErrors(t *testing.T) {
withStdin(t, bytes.NewBufferString("hi"))
var claims ioClaims
r, err := openInput("ca", "-", &claims)
require.NoError(t, err)
r.Close()
_, err = openInput("crt", "-", &claims)
require.EqualError(t, err, `-ca and -crt both set to "-", only one input may read from stdin`)
}
func Test_writeOutput_stdout(t *testing.T) {
out := &bytes.Buffer{}
err := writeOutput("-", []byte("payload"), 0600, out)
require.NoError(t, err)
assert.Equal(t, "payload", out.String())
}
func Test_writeOutput_file(t *testing.T) {
dir := t.TempDir()
p := filepath.Join(dir, "f")
out := &bytes.Buffer{}
err := writeOutput(p, []byte("payload"), 0600, out)
require.NoError(t, err)
assert.Empty(t, out.String())
got, err := os.ReadFile(p)
require.NoError(t, err)
assert.Equal(t, []byte("payload"), got)
}
func Test_reserveOutputs_noConflict(t *testing.T) {
var claims ioClaims
require.NoError(t, reserveOutputs(&claims,
"out-key", "/tmp/key",
"out-crt", "-",
"out-qr", "",
))
assert.Equal(t, "out-crt", claims.out)
}
func Test_reserveOutputs_conflict(t *testing.T) {
var claims ioClaims
err := reserveOutputs(&claims,
"out-key", "-",
"out-crt", "-",
)
require.EqualError(t, err, `-out-key and -out-crt both set to "-", only one output may write to stdout`)
}
func Test_reserveOutputs_panicsOnOddPairs(t *testing.T) {
defer func() {
r := recover()
require.NotNil(t, r)
}()
var claims ioClaims
_ = reserveOutputs(&claims, "out-key")
}
func Test_reserveInputs_noConflict(t *testing.T) {
var claims ioClaims
require.NoError(t, reserveInputs(&claims,
"ca-key", "/tmp/ca.key",
"ca-crt", "-",
"in-pub", "",
))
assert.Equal(t, "ca-crt", claims.in)
}
func Test_reserveInputs_conflict(t *testing.T) {
var claims ioClaims
err := reserveInputs(&claims,
"ca-key", "-",
"ca-crt", "-",
)
require.EqualError(t, err, `-ca-key and -ca-crt both set to "-", only one input may read from stdin`)
}
func Test_claimIn_idempotent(t *testing.T) {
// pre-claim then a lazy re-claim of the same flag should be a no-op
var claims ioClaims
require.NoError(t, claims.claimIn("ca-key"))
require.NoError(t, claims.claimIn("ca-key"))
assert.Equal(t, "ca-key", claims.in)
}
func Test_claimOut_idempotent(t *testing.T) {
var claims ioClaims
require.NoError(t, claims.claimOut("out-crt"))
require.NoError(t, claims.claimOut("out-crt"))
assert.Equal(t, "out-crt", claims.out)
}
func Test_isStdio(t *testing.T) {
assert.True(t, isStdio("-"))
assert.False(t, isStdio(""))
assert.False(t, isStdio("./-"))
assert.False(t, isStdio("foo"))
}
+4 -13
View File
@@ -39,26 +39,18 @@ func verify(args []string, out io.Writer, errOut io.Writer) error {
return err return err
} }
var claims ioClaims caFile, err := os.Open(*vf.caPath)
if err := reserveInputs(&claims,
"ca", *vf.caPath,
"crt", *vf.certPath,
); err != nil {
return err
}
caReader, err := openInput("ca", *vf.caPath, &claims)
if err != nil { if err != nil {
return fmt.Errorf("error while reading ca: %w", err) return fmt.Errorf("error while reading ca: %w", err)
} }
defer caReader.Close() defer caFile.Close()
caPool, err := cert.NewCAPoolFromPEMReader(caReader) caPool, err := cert.NewCAPoolFromPEMReader(caFile)
if err != nil && !errors.Is(err, cert.ErrExpired) { if err != nil && !errors.Is(err, cert.ErrExpired) {
return fmt.Errorf("error while adding ca cert to pool: %w", err) return fmt.Errorf("error while adding ca cert to pool: %w", err)
} }
rawCert, err := readInput("crt", *vf.certPath, &claims) rawCert, err := os.ReadFile(*vf.certPath)
if err != nil { if err != nil {
return fmt.Errorf("unable to read crt: %w", err) return fmt.Errorf("unable to read crt: %w", err)
} }
@@ -93,7 +85,6 @@ func verifySummary() string {
func verifyHelp(out io.Writer) { func verifyHelp(out io.Writer) {
vf := newVerifyFlags() vf := newVerifyFlags()
_, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n")) _, _ = out.Write([]byte("Usage of " + os.Args[0] + " " + verifySummary() + "\n"))
_, _ = out.Write([]byte(stdioHelpText))
vf.set.SetOutput(out) vf.set.SetOutput(out)
vf.set.PrintDefaults() vf.set.PrintDefaults()
} }
-44
View File
@@ -23,7 +23,6 @@ func Test_verifyHelp(t *testing.T) {
assert.Equal( assert.Equal(
t, t,
"Usage of "+os.Args[0]+" verify <flags>: verifies a certificate isn't expired and was signed by a trusted authority.\n"+ "Usage of "+os.Args[0]+" verify <flags>: verifies a certificate isn't expired and was signed by a trusted authority.\n"+
" Pass \"-\" to any path flag to read from stdin or write to stdout.\n"+
" -ca string\n"+ " -ca string\n"+
" \tRequired: path to a file containing one or more ca certificates\n"+ " \tRequired: path to a file containing one or more ca certificates\n"+
" -crt string\n"+ " -crt string\n"+
@@ -123,46 +122,3 @@ func Test_verify(t *testing.T) {
assert.Empty(t, eb.String()) assert.Empty(t, eb.String())
require.NoError(t, err) require.NoError(t, err)
} }
func Test_verify_stdio(t *testing.T) {
ob := &bytes.Buffer{}
eb := &bytes.Buffer{}
caPub, caPriv, _ := ed25519.GenerateKey(rand.Reader)
ca, _ := NewTestCaCert("test-ca", caPub, caPriv, time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour*2), nil, nil, nil)
caPEM, _ := ca.MarshalPEM()
crt, _ := NewTestCert(ca, caPriv, "test-cert", time.Now().Add(time.Hour*-1), time.Now().Add(time.Hour), nil, nil, nil)
crtPEM, _ := crt.MarshalPEM()
caFile, err := os.CreateTemp("", "verify-ca")
require.NoError(t, err)
defer os.Remove(caFile.Name())
caFile.Write(caPEM)
// crt on stdin, ca on disk
withStdin(t, bytes.NewReader(crtPEM))
require.NoError(t, verify([]string{"-ca", caFile.Name(), "-crt", "-"}, ob, eb))
assert.Empty(t, ob.String())
assert.Empty(t, eb.String())
// ca on stdin, crt on disk
certFile, err := os.CreateTemp("", "verify-cert")
require.NoError(t, err)
defer os.Remove(certFile.Name())
certFile.Write(crtPEM)
withStdin(t, bytes.NewReader(caPEM))
ob.Reset()
eb.Reset()
require.NoError(t, verify([]string{"-ca", "-", "-crt", certFile.Name()}, ob, eb))
assert.Empty(t, ob.String())
assert.Empty(t, eb.String())
// both flags on stdin should error
withStdin(t, bytes.NewReader(caPEM))
ob.Reset()
eb.Reset()
require.EqualError(t, verify([]string{"-ca", "-", "-crt", "-"}, ob, eb),
`-ca and -crt both set to "-", only one input may read from stdin`)
}
+6 -13
View File
@@ -53,12 +53,7 @@ func main() {
l := logging.NewLogger(os.Stdout) l := logging.NewLogger(os.Stdout)
if *serviceFlag != "" { if *serviceFlag != "" {
if *configTest { if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
fmt.Println("-test is not supported with -service, run the config test without -service")
os.Exit(1)
}
if err := doService(configPath, Build, serviceFlag); err != nil {
l.Error("Service command failed", "error", err) l.Error("Service command failed", "error", err)
os.Exit(1) os.Exit(1)
} }
@@ -66,13 +61,10 @@ func main() {
} }
if *configPath == "" { if *configPath == "" {
p, err := config.DefaultPath() fmt.Println("-config flag must be set")
if err != nil { flag.Usage()
fmt.Println(err)
os.Exit(1) os.Exit(1)
} }
*configPath = p
}
c := config.NewC(l) c := config.NewC(l)
err := c.Load(*configPath) err := c.Load(*configPath)
@@ -98,14 +90,15 @@ func main() {
} }
if !*configTest { if !*configTest {
if err := ctrl.Start(); err != nil { wait, err := ctrl.Start()
if err != nil {
util.LogWithContextIfNeeded("Error while running", err, l) util.LogWithContextIfNeeded("Error while running", err, l)
os.Exit(1) os.Exit(1)
} }
go ctrl.ShutdownBlock() go ctrl.ShutdownBlock()
if err := ctrl.Wait(); err != nil { if err := wait(); err != nil {
l.Error("Nebula stopped due to fatal error", "error", err) l.Error("Nebula stopped due to fatal error", "error", err)
os.Exit(2) os.Exit(2)
} }
+20 -26
View File
@@ -4,6 +4,7 @@ import (
"fmt" "fmt"
"log" "log"
"os" "os"
"path/filepath"
"github.com/kardianos/service" "github.com/kardianos/service"
"github.com/slackhq/nebula" "github.com/slackhq/nebula"
@@ -15,6 +16,7 @@ var logger service.Logger
type program struct { type program struct {
configPath *string configPath *string
configTest *bool
build string build string
control *nebula.Control control *nebula.Control
} }
@@ -40,47 +42,39 @@ func (p *program) Start(s service.Service) error {
} }
}) })
p.control, err = nebula.Main(c, false, Build, l, nil) p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
if err != nil { if err != nil {
return err return err
} }
if err := p.control.Start(); err != nil { p.control.Start()
return err
}
// Nebula can stop itself on a fatal packet reader error, make sure to log it if it happens.
go func() {
if err := p.control.Wait(); err != nil {
logger.Error(fmt.Sprintf("Nebula stopped due to fatal error: %v", err))
os.Exit(2)
}
}()
return nil return nil
} }
func (p *program) Stop(s service.Service) error { func (p *program) Stop(s service.Service) error {
logger.Info("Nebula service stopping.") logger.Info("Nebula service stopping.")
if p.control == nil {
return nil
}
p.control.Stop() p.control.Stop()
// block until nebula has fully drained before reporting stopped.
// error logging is handled by Start.
_ = p.control.Wait()
return nil return nil
} }
func doService(configPath *string, build string, serviceFlag *string) error { func fileExists(filename string) bool {
_, err := os.Stat(filename)
if os.IsNotExist(err) {
return false
}
return true
}
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
if *configPath == "" { if *configPath == "" {
p, err := config.DefaultPath() ex, err := os.Executable()
if err != nil { if err != nil {
return err return err
} }
*configPath = p *configPath = filepath.Dir(ex) + "/config.yaml"
if !fileExists(*configPath) {
*configPath = filepath.Dir(ex) + "/config.yml"
}
} }
svcConfig := &service.Config{ svcConfig := &service.Config{
@@ -92,6 +86,7 @@ func doService(configPath *string, build string, serviceFlag *string) error {
prg := &program{ prg := &program{
configPath: configPath, configPath: configPath,
configTest: configTest,
build: build, build: build,
} }
@@ -123,9 +118,8 @@ func doService(configPath *string, build string, serviceFlag *string) error {
switch *serviceFlag { switch *serviceFlag {
case "run": case "run":
if err := s.Run(); err != nil { if err := s.Run(); err != nil {
// Route any errors to the system logger and report the failure // Route any errors to the system logger
logger.Error(err) logger.Error(err)
return err
} }
default: default:
if err := service.Control(s, *serviceFlag); err != nil { if err := service.Control(s, *serviceFlag); err != nil {
+5 -7
View File
@@ -50,13 +50,10 @@ func main() {
} }
if *configPath == "" { if *configPath == "" {
p, err := config.DefaultPath() fmt.Println("-config flag must be set")
if err != nil { flag.Usage()
fmt.Println(err)
os.Exit(1) os.Exit(1)
} }
*configPath = p
}
l := logging.NewLogger(os.Stdout) l := logging.NewLogger(os.Stdout)
@@ -84,7 +81,8 @@ func main() {
} }
if !*configTest { if !*configTest {
if err := ctrl.Start(); err != nil { wait, err := ctrl.Start()
if err != nil {
util.LogWithContextIfNeeded("Error while running", err, l) util.LogWithContextIfNeeded("Error while running", err, l)
os.Exit(1) os.Exit(1)
} }
@@ -92,7 +90,7 @@ func main() {
go ctrl.ShutdownBlock() go ctrl.ShutdownBlock()
notifyReady(l) notifyReady(l)
if err := ctrl.Wait(); err != nil { if err := wait(); err != nil {
l.Error("Nebula stopped due to fatal error", "error", err) l.Error("Nebula stopped due to fatal error", "error", err)
os.Exit(2) os.Exit(2)
} }
-29
View File
@@ -1,29 +0,0 @@
package config
import (
"fmt"
"os"
"path/filepath"
)
// DefaultPath returns a path to a config file alongside the running executable, preferring config.yaml over config.yml.
// If neither file exists an error is returned that names both paths checked.
func DefaultPath() (string, error) {
ex, err := os.Executable()
if err != nil {
return "", err
}
return defaultPathInDir(filepath.Dir(ex))
}
func defaultPathInDir(dir string) (string, error) {
yamlPath := filepath.Join(dir, "config.yaml")
if _, err := os.Stat(yamlPath); err == nil {
return yamlPath, nil
}
ymlPath := filepath.Join(dir, "config.yml")
if _, err := os.Stat(ymlPath); err == nil {
return ymlPath, nil
}
return "", fmt.Errorf("no default config found at %s or %s", yamlPath, ymlPath)
}
-67
View File
@@ -1,67 +0,0 @@
package config
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestDefaultPathInDir(t *testing.T) {
t.Run("prefers config.yaml when both exist", func(t *testing.T) {
dir := t.TempDir()
want := filepath.Join(dir, "config.yaml")
other := filepath.Join(dir, "config.yml")
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
require.NoError(t, os.WriteFile(other, []byte("a: 2"), 0644))
got, err := defaultPathInDir(dir)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("returns config.yaml when only it exists", func(t *testing.T) {
dir := t.TempDir()
want := filepath.Join(dir, "config.yaml")
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
got, err := defaultPathInDir(dir)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("falls back to config.yml when only it exists", func(t *testing.T) {
dir := t.TempDir()
want := filepath.Join(dir, "config.yml")
require.NoError(t, os.WriteFile(want, []byte("a: 1"), 0644))
got, err := defaultPathInDir(dir)
require.NoError(t, err)
assert.Equal(t, want, got)
})
t.Run("errors when neither exists and names both paths", func(t *testing.T) {
dir := t.TempDir()
got, err := defaultPathInDir(dir)
assert.Empty(t, got)
require.Error(t, err)
assert.Contains(t, err.Error(), filepath.Join(dir, "config.yaml"))
assert.Contains(t, err.Error(), filepath.Join(dir, "config.yml"))
})
}
func TestDefaultPath(t *testing.T) {
got, err := DefaultPath()
if err != nil {
ex, exErr := os.Executable()
require.NoError(t, exErr)
assert.Contains(t, err.Error(), filepath.Dir(ex))
return
}
ex, err := os.Executable()
require.NoError(t, err)
assert.Equal(t, filepath.Dir(ex), filepath.Dir(got))
assert.Contains(t, []string{"config.yaml", "config.yml"}, filepath.Base(got))
}
+10 -2
View File
@@ -136,6 +136,14 @@ func (cm *connectionManager) getAndResetTrafficCheck(h *HostInfo, now time.Time)
return in, out 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) { func (cm *connectionManager) Start(ctx context.Context) {
clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration) clockSource := time.NewTicker(cm.trafficTimer.t.tickDuration)
defer clockSource.Stop() defer clockSource.Stop()
@@ -298,8 +306,8 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
} else { } else {
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu)) cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
cm.l.Info("send CreateRelayRequest", cm.l.Info("send CreateRelayRequest",
"relayFrom", relayFrom, "relayFrom", req.RelayFromAddr,
"relayTo", relayTo, "relayTo", req.RelayToAddr,
"initiatorRelayIndex", req.InitiatorRelayIndex, "initiatorRelayIndex", req.InitiatorRelayIndex,
"responderRelayIndex", req.ResponderRelayIndex, "responderRelayIndex", req.ResponderRelayIndex,
"vpnAddrs", newhostinfo.vpnAddrs, "vpnAddrs", newhostinfo.vpnAddrs,
+25 -51
View File
@@ -69,29 +69,29 @@ type ControlHostInfo struct {
} }
// Start actually runs nebula, this is a nonblocking call. // Start actually runs nebula, this is a nonblocking call.
// Use Wait to block until nebula has fully stopped and to learn whether a fatal reader error caused the shutdown. // The returned function blocks until nebula has fully stopped and returns the
func (c *Control) Start() error { // first fatal reader error (if any). A nil error means nebula shut down
// gracefully; a non-nil error means a reader hit an unexpected failure that
// triggered the shutdown.
func (c *Control) Start() (func() error, error) {
c.stateLock.Lock() c.stateLock.Lock()
defer c.stateLock.Unlock() defer c.stateLock.Unlock()
switch c.state { switch c.state {
case StateReady: case StateReady:
//yay! //yay!
case StateStopped, StateStopping: case StateStopped, StateStopping:
return ErrAlreadyStopped return nil, ErrAlreadyStopped
case StateStarted: case StateStarted:
return ErrAlreadyStarted return nil, ErrAlreadyStarted
default: default:
return ErrUnknownState return nil, ErrUnknownState
} }
// Activate the interface // Activate the interface
err := c.f.activate() err := c.f.activate()
if err != nil { if err != nil {
// Cancel before Close so a caller returning from Wait always observes a dead Context
c.cancel()
_ = c.f.Close()
c.state = StateStopped c.state = StateStopped
return err return nil, err
} }
// Call all the delayed funcs that waited patiently for the interface to be created. // Call all the delayed funcs that waited patiently for the interface to be created.
@@ -114,9 +114,13 @@ func (c *Control) Start() error {
c.f.triggerShutdown = c.Stop c.f.triggerShutdown = c.Stop
// Start reading packets. // Start reading packets.
c.f.run() out, err := c.f.run()
if err != nil {
c.state = StateStopped
return nil, err
}
c.state = StateStarted c.state = StateStarted
return nil return out, nil
} }
func (c *Control) State() RunState { func (c *Control) State() RunState {
@@ -129,26 +133,10 @@ func (c *Control) Context() context.Context {
return c.ctx return c.ctx
} }
// Stop tears nebula down, closing all tunnels and releasing everything it holds. // Stop is a non-blocking call that signals nebula to close all tunnels and shut down
// Use Wait to block until the shutdown has completed.
// A Control that has been stopped cannot be started again, Start will return ErrAlreadyStopped.
func (c *Control) Stop() { func (c *Control) Stop() {
c.stateLock.Lock() c.stateLock.Lock()
switch c.state { if c.state != StateStarted {
case StateStarted:
// Fall through to the full teardown below
case StateReady:
// Never started
c.cancel()
c.state = StateStopped
if err := c.f.Close(); err != nil {
c.l.Error("Close interface failed", "error", err)
}
c.stateLock.Unlock()
return
default:
c.stateLock.Unlock() c.stateLock.Unlock()
// We are stopping or stopped already // We are stopping or stopped already
return return
@@ -157,26 +145,19 @@ func (c *Control) Stop() {
c.state = StateStopping c.state = StateStopping
c.stateLock.Unlock() c.stateLock.Unlock()
// Closing tunnels can be slow with a large hostmap, don't hold the lock for it // Stop the handshakeManager (and other services), to prevent new tunnels from
// being created while we're shutting them all down.
c.cancel() c.cancel()
c.CloseAllTunnels(false)
c.stateLock.Lock() c.CloseAllTunnels(false)
c.state = StateStopped
if err := c.f.Close(); err != nil { if err := c.f.Close(); err != nil {
c.l.Error("Close interface failed", "error", err) c.l.Error("Close interface failed", "error", err)
} }
c.stateLock.Lock()
c.state = StateStopped
c.stateLock.Unlock() c.stateLock.Unlock()
} }
// Wait blocks until nebula has fully stopped, either via Stop or an internal fatal error,
// and returns the first fatal packet reader error if there was one.
// It is safe to call from multiple goroutines and at any point in the lifecycle,
// but a Wait on a Control that is never started and never stopped will block forever.
func (c *Control) Wait() error {
return c.f.wait()
}
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled // ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
func (c *Control) ShutdownBlock() { func (c *Control) ShutdownBlock() {
sigChan := make(chan os.Signal, 1) sigChan := make(chan os.Signal, 1)
@@ -189,15 +170,8 @@ func (c *Control) ShutdownBlock() {
c.Stop() c.Stop()
} }
// RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change. // RebindUDPServer asks the UDP listener to rebind it's listener. Mainly used on mobile clients when interfaces change
func (c *Control) RebindUDPServer() { func (c *Control) RebindUDPServer() {
c.stateLock.Lock()
defer c.stateLock.Unlock()
if c.state != StateStarted {
return
}
_ = c.f.outside.Rebind() _ = c.f.outside.Rebind()
// Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0 // Trigger a lighthouse update, useful for mobile clients that should have an update interval of 0
@@ -331,7 +305,7 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
c.l.Debug("Sending close tunnel message", c.l.Debug("Sending close tunnel message",
"vpnAddrs", h.vpnAddrs, "vpnAddrs", h.vpnAddrs,
"udpAddr", h.GetRemote(), "udpAddr", h.remote,
) )
closed++ closed++
} }
@@ -376,7 +350,7 @@ func copyHostInfo(h *HostInfo, preferredRanges []netip.Prefix) ControlHostInfo {
RemoteAddrs: h.remotes.CopyAddrs(preferredRanges), RemoteAddrs: h.remotes.CopyAddrs(preferredRanges),
CurrentRelaysToMe: h.relayState.CopyRelayIps(), CurrentRelaysToMe: h.relayState.CopyRelayIps(),
CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(), CurrentRelaysThroughMe: h.relayState.CopyRelayForIps(),
CurrentRemote: h.GetRemote(), CurrentRemote: h.remote,
} }
for i, a := range h.vpnAddrs { for i, a := range h.vpnAddrs {
-309
View File
@@ -1,309 +0,0 @@
package nebula
import (
"context"
"errors"
"io"
"net/netip"
"sync"
"testing"
"time"
"github.com/gaissmai/bart"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/overlay/batch"
"github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/routing"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
type fakeDevice struct {
closeOnce sync.Once
closedCh chan struct{}
closed bool
}
func newFakeDevice() *fakeDevice {
return &fakeDevice{closedCh: make(chan struct{})}
}
// Read blocks until Close like a real tun with no traffic, then reports EOF
// the same way a closed device does
func (d *fakeDevice) Read() ([]tio.Packet, error) {
<-d.closedCh
return nil, io.EOF
}
func (d *fakeDevice) Write(p []byte) (int, error) { return len(p), nil }
func (d *fakeDevice) Close() error {
d.closeOnce.Do(func() {
d.closed = true
close(d.closedCh)
})
return nil
}
func (d *fakeDevice) Activate() error { return nil }
func (d *fakeDevice) Networks() []netip.Prefix { return nil }
func (d *fakeDevice) Name() string { return "fake" }
func (d *fakeDevice) RoutesFor(netip.Addr) routing.Gateways { return nil }
func (d *fakeDevice) Queues(int) ([]tio.Queue, error) { return []tio.Queue{d}, nil }
// newReadyControl hand-builds the minimum Control that Main would have
// produced right before Start, including the construction token NewInterface
// takes so waiters block until Close releases the resources
func newReadyControl(t *testing.T) (*Control, *fakeDevice, *fakeConn) {
l := test.NewLogger()
dev := newFakeDevice()
conn := &fakeConn{}
ctx, cancel := context.WithCancel(context.Background())
myVpnNet := netip.MustParsePrefix("10.128.0.1/16")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(ctx, l, config.NewC(l), cs, nil, nil)
require.NoError(t, err)
f := &Interface{
ctx: ctx,
inside: dev,
outside: conn,
writers: []udp.Conn{conn},
batchers: make([]batch.RxBatcher, 1),
routines: 1,
hostMap: newHostMap(l),
lightHouse: lh,
l: l,
}
f.wg.Add(1)
return &Control{
state: StateReady,
f: f,
l: l,
ctx: ctx,
cancel: cancel,
}, dev, conn
}
func TestControl_StopBeforeStart(t *testing.T) {
c, dev, conn := newReadyControl(t)
// A Stop on a never started control must release everything Main acquired
c.Stop()
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed, "the tun device should have been closed")
assert.True(t, conn.closed, "the udp socket should have been closed")
require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled")
// Wait must return promptly now that the resources are released
require.NoError(t, c.Wait())
// A stopped control can never be started
err := c.Start()
require.ErrorIs(t, err, ErrAlreadyStopped)
// A second Stop is a harmless no-op
c.Stop()
assert.Equal(t, StateStopped, c.State())
require.NoError(t, c.Wait())
}
func TestControl_WaitBlocksUntilStop(t *testing.T) {
c, _, _ := newReadyControl(t)
done := make(chan error, 1)
go func() { done <- c.Wait() }()
select {
case <-done:
t.Fatal("Wait returned before Stop")
case <-time.After(50 * time.Millisecond):
}
c.Stop()
select {
case err := <-done:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("Wait did not return after Stop")
}
}
type fakeConn struct {
closed bool
rebinds int
}
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
func (c *fakeConn) ListenOut(_ udp.EncReader, _ func()) error { return nil }
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
func (c *fakeConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
return nil
}
func (c *fakeConn) ReloadConfig(_ *config.C) {}
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
func (c *fakeConn) Close() error { c.closed = true; return nil }
type multiqueueDevice struct {
*fakeDevice
}
// Queues claims multiqueue support but fails to open the second queue,
// exercising the activation error path.
func (d *multiqueueDevice) Queues(n int) ([]tio.Queue, error) {
if n > 1 {
return nil, errors.New("second queue failed to open")
}
return d.fakeDevice.Queues(n)
}
func TestControl_StartMultiqueueFailureReleases(t *testing.T) {
dev := &multiqueueDevice{fakeDevice: newFakeDevice()}
conn := &fakeConn{}
ctx, cancel := context.WithCancel(context.Background())
f := &Interface{
ctx: ctx,
inside: dev,
outside: conn,
writers: []udp.Conn{conn},
batchers: make([]batch.RxBatcher, 2),
routines: 2,
l: test.NewLogger(),
}
f.wg.Add(1)
c := &Control{
state: StateReady,
f: f,
l: test.NewLogger(),
ctx: ctx,
cancel: cancel,
}
// The second reader fails to open, everything must be released
err := c.Start()
require.Error(t, err)
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed, "the tun device should have been closed")
assert.True(t, conn.closed, "the udp socket should have been closed")
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
// And Wait must not hang on the construction token
require.NoError(t, c.Wait())
}
func TestInterface_CloseIsIdempotent(t *testing.T) {
dev := newFakeDevice()
f := &Interface{
inside: dev,
l: test.NewLogger(),
}
f.wg.Add(1)
require.NoError(t, f.Close())
assert.True(t, dev.closed)
// A second Close must not double release the wg token or the device
require.NoError(t, f.Close())
require.NoError(t, f.wait())
}
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
c, dev, conn := newReadyControl(t)
// Mirror what Start wires up, without needing real packet readers
c.f.triggerShutdown = c.Stop
c.state = StateStarted
boom := errors.New("boom")
c.f.onFatal(boom)
require.ErrorIs(t, c.Wait(), boom)
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed)
assert.True(t, conn.closed)
// A second fatal error must not fire the shutdown again or replace the first
c.f.onFatal(errors.New("later"))
require.ErrorIs(t, c.Wait(), boom)
// Wait stays factual, a Stop after the death does not mask the error
c.Stop()
require.ErrorIs(t, c.Wait(), boom)
}
func TestControl_ConcurrentStopAndStart(t *testing.T) {
c, _, _ := newReadyControl(t)
var wg sync.WaitGroup
for i := 0; i < 2; i++ {
wg.Go(func() { c.Stop() })
}
wg.Go(func() { _ = c.Start() })
wg.Go(func() {
_ = c.Wait()
// A returned Wait must always observe the final state, no matter how
// the race resolved
assert.Equal(t, StateStopped, c.State())
})
wg.Wait()
// However the race resolves, the control must end fully stopped with no
// panic and Wait must observe the final state
require.NoError(t, c.Wait())
assert.Equal(t, StateStopped, c.State())
err := c.Start()
require.ErrorIs(t, err, ErrAlreadyStopped)
}
func TestControl_StartStopLifecycle(t *testing.T) {
c, dev, conn := newReadyControl(t)
err := c.Start()
require.NoError(t, err)
assert.Equal(t, StateStarted, c.State())
err = c.Start()
require.ErrorIs(t, err, ErrAlreadyStarted)
// Stop must unpark the reader blocked in the device and release everything
c.Stop()
assert.Equal(t, StateStopped, c.State())
assert.True(t, dev.closed, "the tun device should have been closed")
assert.True(t, conn.closed, "the udp socket should have been closed")
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
// The reader drained off a closed device, that is not a fatal error
require.NoError(t, c.Wait())
err = c.Start()
require.ErrorIs(t, err, ErrAlreadyStopped)
}
func TestControl_RebindIsGatedByState(t *testing.T) {
c, _, conn := newReadyControl(t)
// A rebind before Start reaches nothing, the interface is not up
c.RebindUDPServer()
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
err := c.Start()
require.NoError(t, err)
c.RebindUDPServer()
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
// A rebind racing a completed stop must not touch the closed conn
c.Stop()
require.NoError(t, c.Wait())
c.RebindUDPServer()
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
}
+6 -161
View File
@@ -1,8 +1,6 @@
package nebula package nebula
import ( import (
"bytes"
"log/slog"
"net" "net"
"net/netip" "net/netip"
"reflect" "reflect"
@@ -11,7 +9,6 @@ import (
"github.com/slackhq/nebula/cert" "github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/test" "github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
) )
func TestControl_GetHostInfoByVpnIp(t *testing.T) { func TestControl_GetHostInfoByVpnIp(t *testing.T) {
@@ -45,7 +42,8 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
assert.True(t, ok) assert.True(t, ok)
crt := &dummyCert{} crt := &dummyCert{}
hi := &HostInfo{ hm.unlockedAddHostInfo(&HostInfo{
remote: remote1,
remotes: remotes, remotes: remotes,
ConnectionState: &ConnectionState{ ConnectionState: &ConnectionState{
peerCert: &cert.CachedCertificate{Certificate: crt}, peerCert: &cert.CachedCertificate{Certificate: crt},
@@ -58,14 +56,13 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
relayForByAddr: map[netip.Addr]*Relay{}, relayForByAddr: map[netip.Addr]*Relay{},
relayForByIdx: map[uint32]*Relay{}, relayForByIdx: map[uint32]*Relay{},
}, },
} }, &Interface{})
hi.remote.Store(&remote1)
hm.unlockedAddHostInfo(hi, &Interface{})
vpnIp2, ok := netip.AddrFromSlice(ipNet2.IP) vpnIp2, ok := netip.AddrFromSlice(ipNet2.IP)
assert.True(t, ok) assert.True(t, ok)
hi2 := &HostInfo{ hm.unlockedAddHostInfo(&HostInfo{
remote: remote1,
remotes: remotes, remotes: remotes,
ConnectionState: &ConnectionState{ ConnectionState: &ConnectionState{
peerCert: nil, peerCert: nil,
@@ -78,9 +75,7 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
relayForByAddr: map[netip.Addr]*Relay{}, relayForByAddr: map[netip.Addr]*Relay{},
relayForByIdx: map[uint32]*Relay{}, relayForByIdx: map[uint32]*Relay{},
}, },
} }, &Interface{})
hi2.remote.Store(&remote1)
hm.unlockedAddHostInfo(hi2, &Interface{})
c := Control{ c := Control{
state: StateReady, state: StateReady,
@@ -124,153 +119,3 @@ func assertFields(t *testing.T, expected []string, actualStruct any) {
assert.Equal(t, expected, fields) 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)
}
})
}
}
-8
View File
@@ -125,14 +125,6 @@ func (c *Control) GetHostmap() *HostMap {
return c.f.hostMap return c.f.hostMap
} }
// GetHostmapIndexCount returns the number of entries in the main hostmap Indexes table, holding
// the hostmap read lock so tests can poll it while connection manager churns tunnels.
func (c *Control) GetHostmapIndexCount() int {
c.f.hostMap.RLock()
defer c.f.hostMap.RUnlock()
return len(c.f.hostMap.Indexes)
}
func (c *Control) GetF() *Interface { func (c *Control) GetF() *Interface {
return c.f return c.f
} }
+1 -1
View File
@@ -62,7 +62,7 @@ function nebula.dissector(tvbuf, pktinfo, root)
tree:add(pf_version, tvbuf:range(0,1)) tree:add(pf_version, tvbuf:range(0,1))
local type = tree:add(pf_type, tvbuf:range(0,1)) local type = tree:add(pf_type, tvbuf:range(0,1))
local nebula_type = bit.band(tvbuf:range(0,1):uint(), 0x0F) local nebula_type = bit32.band(tvbuf:range(0,1):uint(), 0x0F)
if nebula_type == 0 then if nebula_type == 0 then
local stage = tvbuf(8,8):uint64() local stage = tvbuf(8,8):uint64()
tree:add(pf_subtype_handshake, tvbuf:range(1,1)) tree:add(pf_subtype_handshake, tvbuf:range(1,1))
+16 -82
View File
@@ -11,6 +11,7 @@ import (
"sync" "sync"
"sync/atomic" "sync/atomic"
"github.com/gaissmai/bart"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/slackhq/nebula/config" "github.com/slackhq/nebula/config"
) )
@@ -22,10 +23,7 @@ type dnsServer struct {
dnsMap4 map[string]netip.Addr dnsMap4 map[string]netip.Addr
dnsMap6 map[string]netip.Addr dnsMap6 map[string]netip.Addr
hostMap *HostMap hostMap *HostMap
pki *PKI myVpnAddrsTable *bart.Lite
// selfHost is the cached FQDN we last seeded for ourselves
selfHost string
mux *dns.ServeMux mux *dns.ServeMux
@@ -57,14 +55,14 @@ type dnsServer struct {
// they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel // they no-op when DNS isn't enabled. Each Start invocation owns a ctx-cancel
// watcher that tears the listener down on nebula shutdown. The returned // watcher that tears the listener down on nebula shutdown. The returned
// pointer is always non-nil, even on error. // pointer is always non-nil, even on error.
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostMap *HostMap, c *config.C) (*dnsServer, error) { func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, cs *CertState, hostMap *HostMap, c *config.C) (*dnsServer, error) {
ds := &dnsServer{ ds := &dnsServer{
l: l, l: l,
ctx: ctx, ctx: ctx,
dnsMap4: make(map[string]netip.Addr), dnsMap4: make(map[string]netip.Addr),
dnsMap6: make(map[string]netip.Addr), dnsMap6: make(map[string]netip.Addr),
hostMap: hostMap, hostMap: hostMap,
pki: pki, myVpnAddrsTable: cs.myVpnAddrsTable,
} }
ds.mux = dns.NewServeMux() ds.mux = dns.NewServeMux()
ds.mux.HandleFunc(".", ds.handleDnsRequest) ds.mux.HandleFunc(".", ds.handleDnsRequest)
@@ -78,7 +76,6 @@ func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, pki *PKI, hostM
if err := ds.reload(c, true); err != nil { if err := ds.reload(c, true); err != nil {
return ds, err return ds, err
} }
ds.seedSelf()
return ds, nil return ds, nil
} }
@@ -116,7 +113,7 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
d.Stop() d.Stop()
} }
// Drop any records that accumulated while enabled; a later re-enable // Drop any records that accumulated while enabled; a later re-enable
// will repopulate from fresh handshakes and a fresh seedSelf. // will repopulate from fresh handshakes.
d.clearRecords() d.clearRecords()
return nil return nil
} }
@@ -124,14 +121,17 @@ func (d *dnsServer) reload(c *config.C, initial bool) error {
if running == nil { if running == nil {
// Was disabled (or never started); bring it up now. // Was disabled (or never started); bring it up now.
go d.Start() go d.Start()
} else if !sameAddr { return nil
d.shutdownServer(running, runningStarted, "reload")
// Old Start goroutine has now exited; bring up a fresh listener on the new address.
go d.Start()
} }
// Refresh the self entry every enabled reload so cert renewals that change our name or VPN addresses are picked up. if sameAddr {
d.seedSelf() return nil
}
d.shutdownServer(running, runningStarted, "reload")
// Old Start goroutine has now exited; bring up a fresh listener on the
// new address.
go d.Start()
return nil return nil
} }
@@ -249,20 +249,6 @@ func (d *dnsServer) QueryCert(data string) string {
return "" return ""
} }
// The hostmap only ever contains peers we have handshaked with, so it never carries an entry for ourselves.
// Answer self lookups straight from the local cert state.
if cs := d.certState(); cs != nil && cs.myVpnAddrsTable != nil && cs.myVpnAddrsTable.Contains(ip) {
c := cs.GetDefaultCertificate()
if c == nil {
return ""
}
b, err := c.MarshalJSON()
if err != nil {
return ""
}
return string(b)
}
hostinfo := d.hostMap.QueryVpnAddr(ip) hostinfo := d.hostMap.QueryVpnAddr(ip)
if hostinfo == nil { if hostinfo == nil {
return "" return ""
@@ -280,60 +266,12 @@ func (d *dnsServer) QueryCert(data string) string {
return string(b) return string(b)
} }
// clearRecords drops all DNS records, including the self entry. // clearRecords drops all DNS records.
func (d *dnsServer) clearRecords() { func (d *dnsServer) clearRecords() {
d.Lock() d.Lock()
defer d.Unlock() defer d.Unlock()
clear(d.dnsMap4) clear(d.dnsMap4)
clear(d.dnsMap6) clear(d.dnsMap6)
d.selfHost = ""
}
// seedSelf inserts (or refreshes) a record for our own cert name pointing at our VPN addresses,
// so a single-lighthouse network can resolve the lighthouse's own hostname without the two-process workaround.
func (d *dnsServer) seedSelf() {
if !d.enabled.Load() {
return
}
cs := d.certState()
if cs == nil {
return
}
c := cs.GetDefaultCertificate()
if c == nil {
return
}
newHost := strings.ToLower(c.Name()) + "."
d.Lock()
defer d.Unlock()
if d.selfHost != "" && d.selfHost != newHost {
delete(d.dnsMap4, d.selfHost)
delete(d.dnsMap6, d.selfHost)
}
d.selfHost = newHost
delete(d.dnsMap4, newHost)
delete(d.dnsMap6, newHost)
haveV4, haveV6 := false, false
for _, addr := range cs.myVpnAddrs {
if addr.Is4() && !haveV4 {
d.dnsMap4[newHost] = addr
haveV4 = true
} else if addr.Is6() && !haveV6 {
d.dnsMap6[newHost] = addr
haveV6 = true
}
if haveV4 && haveV6 {
break
}
}
}
func (d *dnsServer) certState() *CertState {
if d.pki == nil {
return nil
}
return d.pki.getCertState()
} }
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host` // Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
@@ -371,12 +309,8 @@ func (d *dnsServer) isSelfNebulaOrLocalhost(addr string) bool {
return true return true
} }
cs := d.certState()
if cs == nil || cs.myVpnAddrsTable == nil {
return false
}
//if we found it in this table, it's good //if we found it in this table, it's good
return cs.myVpnAddrsTable.Contains(b) return d.myVpnAddrsTable.Contains(b)
} }
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) { func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
-89
View File
@@ -9,10 +9,7 @@ import (
"testing" "testing"
"time" "time"
"github.com/gaissmai/bart"
"github.com/miekg/dns" "github.com/miekg/dns"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/cert_test"
"github.com/slackhq/nebula/config" "github.com/slackhq/nebula/config"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -279,92 +276,6 @@ func TestDnsServer_Stop_beforeBind_doesNotHang(t *testing.T) {
} }
} }
// newTestPKI builds a minimal *PKI with a single v1 cert whose name and
// VPN addresses are caller-provided, suitable for exercising seedSelf and
// QueryCert self handling.
func newTestPKI(t *testing.T, name string, addrs []netip.Addr) *PKI {
t.Helper()
networks := make([]netip.Prefix, 0, len(addrs))
for _, a := range addrs {
bits := 32
if a.Is6() {
bits = 128
}
networks = append(networks, netip.PrefixFrom(a, bits))
}
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil)
c, _, _, _ := cert_test.NewTestCert(cert.Version2, cert.Curve_CURVE25519, ca, caKey, name, time.Time{}, time.Time{}, networks, nil, nil)
addrsTable := new(bart.Lite)
for _, a := range addrs {
addrsTable.Insert(netip.PrefixFrom(a, a.BitLen()))
}
cs := &CertState{
v2Cert: c,
initiatingVersion: cert.Version2,
myVpnAddrs: addrs,
myVpnAddrsTable: addrsTable,
}
pki := &PKI{}
pki.cs.Store(cs)
return pki
}
func TestDnsServer_seedSelf_addsOwnRecord(t *testing.T) {
ds, c := newTestDnsServer(t)
myV4 := netip.MustParseAddr("10.0.0.1")
myV6 := netip.MustParseAddr("fd00::1")
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{myV4, myV6})
setDnsConfig(c, "127.0.0.1", "0", true, true)
require.NoError(t, ds.reload(c, true))
ds.seedSelf()
got4, exists := ds.Query(dns.TypeA, "lighthouse.")
assert.True(t, exists)
assert.Equal(t, myV4, got4)
got6, exists := ds.Query(dns.TypeAAAA, "lighthouse.")
assert.True(t, exists)
assert.Equal(t, myV6, got6)
}
func TestDnsServer_seedSelf_disabled_noOp(t *testing.T) {
ds, c := newTestDnsServer(t)
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
setDnsConfig(c, "127.0.0.1", "0", true, false)
require.NoError(t, ds.reload(c, true))
ds.seedSelf()
_, exists := ds.Query(dns.TypeA, "lighthouse.")
assert.False(t, exists)
}
func TestDnsServer_clearRecords_dropsSelfHost(t *testing.T) {
ds, c := newTestDnsServer(t)
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{netip.MustParseAddr("10.0.0.1")})
setDnsConfig(c, "127.0.0.1", "0", true, true)
require.NoError(t, ds.reload(c, true))
ds.seedSelf()
require.NotEmpty(t, ds.selfHost)
ds.clearRecords()
assert.Empty(t, ds.selfHost)
_, exists := ds.Query(dns.TypeA, "lighthouse.")
assert.False(t, exists)
}
func TestDnsServer_QueryCert_returnsOwnCert(t *testing.T) {
ds, _ := newTestDnsServer(t)
myV4 := netip.MustParseAddr("10.0.0.1")
ds.pki = newTestPKI(t, "lighthouse", []netip.Addr{myV4})
got := ds.QueryCert(myV4.String() + ".")
assert.NotEmpty(t, got, "TXT lookup of our own VPN address should return our cert")
other := netip.MustParseAddr("10.0.0.99")
assert.Empty(t, ds.QueryCert(other.String()+"."), "unknown peer IP should return nothing")
}
func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) { func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) {
port := freeUDPPort(t) port := freeUDPPort(t)
ds, c := newTestDnsServer(t) ds, c := newTestDnsServer(t)
-10
View File
@@ -1,16 +1,6 @@
FROM gcr.io/distroless/static:latest FROM gcr.io/distroless/static:latest
ARG TARGETOS TARGETARCH 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 /nebula
COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert COPY build/$TARGETOS-$TARGETARCH/nebula-cert /nebula-cert
-85
View File
@@ -1,85 +0,0 @@
//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
}
}
}
+30 -109
View File
@@ -405,7 +405,7 @@ func TestStage1Race(t *testing.T) {
r.Log("Spin until connection manager tears down a tunnel") r.Log("Spin until connection manager tears down a tunnel")
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 { for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet") t.Log("Connection manager hasn't ticked yet")
time.Sleep(time.Second) time.Sleep(time.Second)
@@ -453,11 +453,9 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
r.Log("Nuke my hostmap") r.Log("Nuke my hostmap")
myHostmap := myControl.GetHostmap() myHostmap := myControl.GetHostmap()
myHostmap.Lock()
myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{} myHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
myHostmap.Indexes = map[uint32]*nebula.HostInfo{} myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{} myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
myHostmap.Unlock()
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again"))) myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
p = r.RouteForAllUntilTxTun(theirControl) p = r.RouteForAllUntilTxTun(theirControl)
@@ -467,10 +465,10 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
r.Log("Wait for the dead index to go away") r.Log("Wait for the dead index to go away")
start := theirControl.GetHostmapIndexCount() start := len(theirControl.GetHostmap().Indexes)
for { for {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
if theirControl.GetHostmapIndexCount() < start { if len(theirControl.GetHostmap().Indexes) < start {
break break
} }
time.Sleep(time.Second) time.Sleep(time.Second)
@@ -506,11 +504,9 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
r.Log("Nuke my hostmap") r.Log("Nuke my hostmap")
theirHostmap := theirControl.GetHostmap() theirHostmap := theirControl.GetHostmap()
theirHostmap.Lock()
theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{} theirHostmap.Hosts = map[netip.Addr]*nebula.HostInfo{}
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{} theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{} theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
theirHostmap.Unlock()
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again"))) theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
p = r.RouteForAllUntilTxTun(myControl) p = r.RouteForAllUntilTxTun(myControl)
@@ -521,10 +517,10 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
r.Log("Wait for the dead index to go away") r.Log("Wait for the dead index to go away")
start := myControl.GetHostmapIndexCount() start := len(myControl.GetHostmap().Indexes)
for { for {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
if myControl.GetHostmapIndexCount() < start { if len(myControl.GetHostmap().Indexes) < start {
break break
} }
time.Sleep(time.Second) time.Sleep(time.Second)
@@ -632,10 +628,10 @@ func TestReestablishRelays(t *testing.T) {
r.Log("Close the tunnel") r.Log("Close the tunnel")
relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true) relayControl.CloseTunnel(theirVpnIpNet[0].Addr(), true)
start := myControl.GetHostmapIndexCount() start := len(myControl.GetHostmap().Indexes)
curIndexes := myControl.GetHostmapIndexCount() curIndexes := len(myControl.GetHostmap().Indexes)
for curIndexes >= start { for curIndexes >= start {
curIndexes = myControl.GetHostmapIndexCount() curIndexes = len(myControl.GetHostmap().Indexes)
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes) r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail"))) myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
@@ -823,18 +819,18 @@ func TestStage1RaceRelays2(t *testing.T) {
t.Log("Wait until we remove extra tunnels") t.Log("Wait until we remove extra tunnels")
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d", t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
myControl.GetHostmapIndexCount(), len(myControl.GetHostmap().Indexes),
theirControl.GetHostmapIndexCount(), len(theirControl.GetHostmap().Indexes),
relayControl.GetHostmapIndexCount(), len(relayControl.GetHostmap().Indexes),
) )
hostInfos := myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount() hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
retries := 60 retries := 60
for hostInfos > 6 && retries > 0 { for hostInfos > 6 && retries > 0 {
hostInfos = myControl.GetHostmapIndexCount() + theirControl.GetHostmapIndexCount() + relayControl.GetHostmapIndexCount() hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d", t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
myControl.GetHostmapIndexCount(), len(myControl.GetHostmap().Indexes),
theirControl.GetHostmapIndexCount(), len(theirControl.GetHostmap().Indexes),
relayControl.GetHostmapIndexCount(), len(relayControl.GetHostmap().Indexes),
) )
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet") t.Log("Connection manager hasn't ticked yet")
@@ -928,24 +924,24 @@ func TestRehandshakingRelays(t *testing.T) {
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl) r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
// We should have two hostinfos on all sides // We should have two hostinfos on all sides
for myControl.GetHostmapIndexCount() != 2 { for len(myControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount()) t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works") r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes") r.Log("yupitdoes")
time.Sleep(time.Second) time.Sleep(time.Second)
} }
t.Logf("myControl hostinfos got cleaned up!") t.Logf("myControl hostinfos got cleaned up!")
for theirControl.GetHostmapIndexCount() != 2 { for len(theirControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount()) t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works") r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes") r.Log("yupitdoes")
time.Sleep(time.Second) time.Sleep(time.Second)
} }
t.Logf("theirControl hostinfos got cleaned up!") t.Logf("theirControl hostinfos got cleaned up!")
for relayControl.GetHostmapIndexCount() != 2 { for len(relayControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount()) t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works") r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes") r.Log("yupitdoes")
@@ -1033,24 +1029,24 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl) r.RenderHostmaps("working hostmaps", myControl, relayControl, theirControl)
// We should have two hostinfos on all sides // We should have two hostinfos on all sides
for myControl.GetHostmapIndexCount() != 2 { for len(myControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", myControl.GetHostmapIndexCount()) t.Logf("Waiting for myControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(myControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works") r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes") r.Log("yupitdoes")
time.Sleep(time.Second) time.Sleep(time.Second)
} }
t.Logf("myControl hostinfos got cleaned up!") t.Logf("myControl hostinfos got cleaned up!")
for theirControl.GetHostmapIndexCount() != 2 { for len(theirControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", theirControl.GetHostmapIndexCount()) t.Logf("Waiting for theirControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(theirControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works") r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes") r.Log("yupitdoes")
time.Sleep(time.Second) time.Sleep(time.Second)
} }
t.Logf("theirControl hostinfos got cleaned up!") t.Logf("theirControl hostinfos got cleaned up!")
for relayControl.GetHostmapIndexCount() != 2 { for len(relayControl.GetHostmap().Indexes) != 2 {
t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", relayControl.GetHostmapIndexCount()) t.Logf("Waiting for relayControl hostinfos (%v != 2) to get cleaned up from lack of use...", len(relayControl.GetHostmap().Indexes))
r.Log("Assert the relay tunnel still works") r.Log("Assert the relay tunnel still works")
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r) assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
r.Log("yupitdoes") r.Log("yupitdoes")
@@ -1127,7 +1123,7 @@ func TestRehandshaking(t *testing.T) {
theirConfig.ReloadConfigString(string(rc)) theirConfig.ReloadConfigString(string(rc))
r.Log("Spin until there is only 1 tunnel") r.Log("Spin until there is only 1 tunnel")
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 { for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet") t.Log("Connection manager hasn't ticked yet")
time.Sleep(time.Second) time.Sleep(time.Second)
@@ -1227,7 +1223,7 @@ func TestRehandshakingLoser(t *testing.T) {
myConfig.ReloadConfigString(string(rc)) myConfig.ReloadConfigString(string(rc))
r.Log("Spin until there is only 1 tunnel") r.Log("Spin until there is only 1 tunnel")
for myControl.GetHostmapIndexCount()+theirControl.GetHostmapIndexCount() > 2 { for len(myControl.GetHostmap().Indexes)+len(theirControl.GetHostmap().Indexes) > 2 {
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r) assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
t.Log("Connection manager hasn't ticked yet") t.Log("Connection manager hasn't ticked yet")
time.Sleep(time.Second) time.Sleep(time.Second)
@@ -1539,78 +1535,3 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
myControl.Stop() myControl.Stop()
theirControl.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()
}
+4 -2
View File
@@ -4,13 +4,15 @@
package e2e package e2e
import ( import (
"log/slog" "io"
"net/netip" "net/netip"
"os" "os"
"strings" "strings"
"testing" "testing"
"time" "time"
"log/slog"
"dario.cat/mergo" "dario.cat/mergo"
"github.com/google/gopacket" "github.com/google/gopacket"
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
@@ -380,7 +382,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
func NewTestLogger() *slog.Logger { func NewTestLogger() *slog.Logger {
v := os.Getenv("TEST_LOGS") v := os.Getenv("TEST_LOGS")
if v == "" { if v == "" {
return slog.New(slog.DiscardHandler) return slog.New(slog.NewTextHandler(io.Discard, nil))
} }
level := slog.LevelInfo level := slog.LevelInfo
+6 -101
View File
@@ -15,7 +15,6 @@ import (
"github.com/slackhq/nebula/header" "github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/udp" "github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gopkg.in/yaml.v3" "gopkg.in/yaml.v3"
) )
@@ -43,8 +42,8 @@ func TestDropInactiveTunnels(t *testing.T) {
r.Log("Go inactive and wait for the tunnels to get dropped") r.Log("Go inactive and wait for the tunnels to get dropped")
waitStart := time.Now() waitStart := time.Now()
for { for {
myIndexes := myControl.GetHostmapIndexCount() myIndexes := len(myControl.GetHostmap().Indexes)
theirIndexes := theirControl.GetHostmapIndexCount() theirIndexes := len(theirControl.GetHostmap().Indexes)
if myIndexes == 0 && theirIndexes == 0 { if myIndexes == 0 && theirIndexes == 0 {
break break
} }
@@ -374,100 +373,6 @@ func TestCrossStackRelaysWork(t *testing.T) {
//relayControl.Stop() //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) { func TestCloseTunnelAuthenticated(t *testing.T) {
t.Parallel() t.Parallel()
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{}) ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
@@ -493,8 +398,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
waitStart := time.Now() waitStart := time.Now()
for { for {
myIndexes := myControl.GetHostmapIndexCount() myIndexes := len(myControl.GetHostmap().Indexes)
theirIndexes := theirControl.GetHostmapIndexCount() theirIndexes := len(theirControl.GetHostmap().Indexes)
if myIndexes == 0 && theirIndexes == 0 { if myIndexes == 0 && theirIndexes == 0 {
break break
} }
@@ -548,8 +453,8 @@ func TestCloseTunnelAuthenticated(t *testing.T) {
r.Log("Injected bogus close tunnel. Let's see!") r.Log("Injected bogus close tunnel. Let's see!")
waitStart = time.Now() waitStart = time.Now()
for { for {
myIndexes := myControl.GetHostmapIndexCount() myIndexes := len(myControl.GetHostmap().Indexes)
theirIndexes := theirControl.GetHostmapIndexCount() theirIndexes := len(theirControl.GetHostmap().Indexes)
if myIndexes == 0 { if myIndexes == 0 {
t.Fatal("myIndexes should not be 0") t.Fatal("myIndexes should not be 0")
} }
+2 -65
View File
@@ -1,11 +1,9 @@
package nebula package nebula
import ( import (
"encoding/binary" "io"
"log/slog" "log/slog"
"testing" "testing"
"golang.org/x/net/ipv4"
) )
func TestInnerECN(t *testing.T) { func TestInnerECN(t *testing.T) {
@@ -52,7 +50,7 @@ func v6WithTC(tc byte) []byte {
} }
func TestApplyOuterECN(t *testing.T) { func TestApplyOuterECN(t *testing.T) {
silent := slog.New(slog.DiscardHandler) silent := slog.New(slog.NewTextHandler(io.Discard, nil))
hi := &HostInfo{} hi := &HostInfo{}
// Build a v4 packet helper with a given inner ECN field. // Build a v4 packet helper with a given inner ECN field.
@@ -125,64 +123,3 @@ func TestApplyOuterECN(t *testing.T) {
}) })
} }
} }
// TestApplyOuterECN_IPv4ChecksumStaysValid guards against H1: folding an outer
// CE mark into the inner IPv4 ToS byte must keep the IPv4 header checksum valid.
// The passthrough emit paths write the packet verbatim, so a stale checksum
// turns an underlay congestion mark into packet loss at the receiver.
func TestApplyOuterECN_IPv4ChecksumStaysValid(t *testing.T) {
silent := slog.New(slog.DiscardHandler)
hi := &HostInfo{}
// 20-byte IPv4 header with DSCP=0x88 and inner ECN = ECT(0). Folding CE
// flips only the low two bits of the ToS byte while leaving DSCP intact.
pkt := []byte{
0x45, 0x88 | ecnECT0, 0, 40,
0x1c, 0x46, 0x40, 0x00,
64, 6, 0, 0,
10, 0, 0, 1,
10, 0, 0, 2,
}
// Stamp a correct header checksum before the fold.
binary.BigEndian.PutUint16(pkt[10:12], ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]))
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
t.Fatal("test setup: initial header checksum invalid")
}
applyOuterECN(pkt, ecnCE, hi, silent)
// CE folded in, DSCP preserved.
if got, want := pkt[1], byte(0x88|ecnCE); got != want {
t.Fatalf("ToS after fold = 0x%02x, want 0x%02x", got, want)
}
// The incremental RFC 1624 update must leave the checksum valid and equal
// to a full recompute over the mutated header.
if !ipv4HeaderChecksumValid(pkt[:ipv4.HeaderLen]) {
t.Fatalf("IPv4 header checksum invalid after CE fold: 0x%04x", binary.BigEndian.Uint16(pkt[10:12]))
}
if got, want := binary.BigEndian.Uint16(pkt[10:12]), ipv4HeaderChecksum(pkt[:ipv4.HeaderLen]); got != want {
t.Fatalf("checksum = 0x%04x, full recompute = 0x%04x", got, want)
}
}
// ipv4HeaderChecksum computes the RFC 1071 IPv4 header checksum over hdr,
// treating the checksum field (bytes 10:12) as zero.
func ipv4HeaderChecksum(hdr []byte) uint16 {
var sum uint32
for i := 0; i+1 < len(hdr); i += 2 {
if i == 10 {
continue // checksum field
}
sum += uint32(hdr[i])<<8 | uint32(hdr[i+1])
}
for sum > 0xffff {
sum = (sum >> 16) + (sum & 0xffff)
}
return ^uint16(sum)
}
// ipv4HeaderChecksumValid reports whether the stored checksum matches a fresh
// computation over the header.
func ipv4HeaderChecksumValid(hdr []byte) bool {
return binary.BigEndian.Uint16(hdr[10:12]) == ipv4HeaderChecksum(hdr)
}
+1 -33
View File
@@ -254,28 +254,6 @@ tun:
# Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic # Default MTU for every packet, safe setting is (and the default) 1300 for internet based traffic
mtu: 1300 mtu: 1300
# Linux only. pin_threads pins each tun reader/encrypt OS thread to a single CPU. This keeps every goroutine's
# batched sends flowing through one XPS-selected NIC TX ring, so packets within a flow stay ordered on the wire
# instead of being sprayed across multiple TX rings and reordered. Not reloadable.
#
# When cpu_affinity is unset, nebula picks CPUs that do NOT service any physical NIC's interrupts (read from
# /sys/class/net/*/device/msi_irqs and /proc/irq/*/effective_affinity_list): an encrypt thread pinned onto a core
# that also runs NAPI for a NIC RX queue fights the softirq for the core and collapses throughput for flows hashed
# to that queue. If the NIC's vectors blanket every allowed CPU (many drivers default to one queue per core) the
# avoidance logs and falls back to the old spread; narrow the NIC's queue/IRQ spread (e.g. `ethtool -X <dev>
# equal N`) or set cpu_affinity explicitly to benefit.
#pin_threads: true
# Linux only. cpu_affinity overrides which CPUs the tun reader threads pin to: a list of CPU IDs, one per routine
# (see the top-level `routines` setting). Lists shorter than `routines` are modulo-cycled across the queues; extra
# entries are ignored. IDs must be within the process's allowed CPU set, so this respects taskset / cgroup cpusets;
# a non-integer or not-allowed entry disables the override and falls back to spreading queues across the allowed
# CPUs. Setting this disables the automatic NIC-IRQ avoidance described under pin_threads — prefer CPUs that don't
# service your underlay NIC's RX queue IRQs. Only meaningful while pin_threads is true. Not reloadable.
#cpu_affinity:
# - 2
# - 4
# Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here # Route based MTU overrides, you have known vpn ip paths that can support larger MTUs you can increase/decrease them here
routes: routes:
#- mtu: 8800 #- mtu: 8800
@@ -412,16 +390,6 @@ logging:
# This setting is reloadable # This setting is reloadable
#inactivity_timeout: 10m #inactivity_timeout: 10m
# ecn (default true) propagates ECN (Explicit Congestion Notification) across the tunnel per RFC 6040: the inner
# packet's ECN codepoint is copied onto the outer carrier header on encapsulation, and an outer CE ("congestion
# experienced") mark is folded back into the inner header on decapsulation. On linux it additionally stamps
# RTAX_FEATURE_ECN on the routes nebula installs, so the kernel actively negotiates ECN for connections to mesh
# prefixes. Disable this only when an underlay middlebox mangles or clears ECN bits unpredictably.
# This setting is reloadable, BUT flipping it at runtime only updates the datapath (the inner<->outer copy/combine).
# The RTAX_FEATURE_ECN flag on already-installed routes is NOT revisited on reload, so nebula must be restarted for
# the route half of this setting to take effect.
#ecn: true
# Nebula security group configuration # Nebula security group configuration
firewall: firewall:
# Action to take when a packet is not allowed by the firewall rules. # Action to take when a packet is not allowed by the firewall rules.
@@ -429,7 +397,7 @@ firewall:
# `drop` (default): silently drop the packet. # `drop` (default): silently drop the packet.
# `reject`: send a reject reply. # `reject`: send a reject reply.
# - For TCP, this will be a RST "Connection Reset" packet. # - For TCP, this will be a RST "Connection Reset" packet.
# - For other protocols, this will be an ICMP "Destination unreachable: Communication administratively prohibited" packet. # - For other protocols, this will be an ICMP port unreachable packet.
outbound_action: drop outbound_action: drop
inbound_action: drop inbound_action: drop
+72 -70
View File
@@ -44,8 +44,8 @@ type Firewall struct {
InRules *FirewallTable InRules *FirewallTable
OutRules *FirewallTable OutRules *FirewallTable
InboundSendReject bool InSendReject bool
OutboundSendReject bool OutSendReject bool
//TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better //TODO: we should have many more options for TCP, an option for ICMP, and mimic the kernel a bit better
// https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt // https://www.kernel.org/doc/Documentation/networking/nf_conntrack-sysctl.txt
@@ -59,8 +59,7 @@ type Firewall struct {
// assignedNetworks is a list of vpn networks assigned to us in the certificate. // assignedNetworks is a list of vpn networks assigned to us in the certificate.
assignedNetworks []netip.Prefix assignedNetworks []netip.Prefix
// unsafeNetworks is the list of unsafe networks issued to us in the certificate hasUnsafeNetworks bool
unsafeNetworks []netip.Prefix
rules string rules string
rulesVersion uint16 rulesVersion uint16
@@ -81,8 +80,8 @@ type firewallMetrics struct {
type FirewallConntrack struct { type FirewallConntrack struct {
sync.Mutex sync.Mutex
Conns map[firewall.Packet]*conn Conns map[firewall.PacketKey]*conn
TimerWheel *TimerWheel[firewall.Packet] TimerWheel *TimerWheel[firewall.PacketKey]
} }
// FirewallTable is the entry point for a rule, the evaluation order is: // FirewallTable is the entry point for a rule, the evaluation order is:
@@ -159,15 +158,16 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
assignedNetworks = append(assignedNetworks, network) assignedNetworks = append(assignedNetworks, network)
} }
unsafeNetworks := c.UnsafeNetworks() hasUnsafeNetworks := false
for _, n := range unsafeNetworks { for _, n := range c.UnsafeNetworks() {
routableNetworks.Insert(n) routableNetworks.Insert(n)
hasUnsafeNetworks = true
} }
return &Firewall{ return &Firewall{
Conntrack: &FirewallConntrack{ Conntrack: &FirewallConntrack{
Conns: make(map[firewall.Packet]*conn), Conns: make(map[firewall.PacketKey]*conn),
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax), TimerWheel: NewTimerWheel[firewall.PacketKey](tmin, tmax),
}, },
InRules: newFirewallTable(), InRules: newFirewallTable(),
OutRules: newFirewallTable(), OutRules: newFirewallTable(),
@@ -176,7 +176,7 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
DefaultTimeout: defaultTimeout, DefaultTimeout: defaultTimeout,
routableNetworks: routableNetworks, routableNetworks: routableNetworks,
assignedNetworks: assignedNetworks, assignedNetworks: assignedNetworks,
unsafeNetworks: unsafeNetworks, hasUnsafeNetworks: hasUnsafeNetworks,
l: l, l: l,
incomingMetrics: firewallMetrics{ incomingMetrics: firewallMetrics{
@@ -216,23 +216,23 @@ func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewal
inboundAction := c.GetString("firewall.inbound_action", "drop") inboundAction := c.GetString("firewall.inbound_action", "drop")
switch inboundAction { switch inboundAction {
case "reject": case "reject":
fw.InboundSendReject = true fw.InSendReject = true
case "drop": case "drop":
fw.InboundSendReject = false fw.InSendReject = false
default: default:
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction) l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
fw.InboundSendReject = false fw.InSendReject = false
} }
outboundAction := c.GetString("firewall.outbound_action", "drop") outboundAction := c.GetString("firewall.outbound_action", "drop")
switch outboundAction { switch outboundAction {
case "reject": case "reject":
fw.OutboundSendReject = true fw.OutSendReject = true
case "drop": case "drop":
fw.OutboundSendReject = false fw.OutSendReject = false
default: default:
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction) l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
fw.OutboundSendReject = false fw.OutSendReject = false
} }
err := AddFirewallRulesFromConfig(l, false, c, fw) err := AddFirewallRulesFromConfig(l, false, c, fw)
@@ -422,7 +422,27 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table")
// Drop returns an error if the packet should be dropped, explaining why. It // Drop returns an error if the packet should be dropped, explaining why. It
// returns nil if the packet should not be dropped. // 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 { //
// key is the dense conntrack key — used as-is for the inConns fast path
// without touching fp at all. fp is the rich Packet form rule matching
// needs (CIDR lookups, family checks); on the conntrack-miss slow path
// Drop ensures fp is hydrated from key (idempotent if the caller already
// filled fp). On accept-via-conntrack the caller's fp is left untouched.
func (f *Firewall) Drop(key firewall.PacketKey, 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.
// Hot path: only the dense key is touched.
if f.inConns(key, h, caPool, localCache) {
return nil
}
// Conntrack miss → rule matching needs the rich Packet form. Hydrate
// from the key if the caller passed a zero-valued fp (the inbound path
// after batch.ParsePacket). Outbound callers Hydrate themselves and
// skip this hop.
if !fp.LocalAddr.IsValid() {
key.Hydrate(fp)
}
// Make sure remote address matches nebula certificate, and determine how to treat it // Make sure remote address matches nebula certificate, and determine how to treat it
if h.networks == nil { if h.networks == nil {
// Simple case: Certificate has one address and no unsafe networks // Simple case: Certificate has one address and no unsafe networks
@@ -456,24 +476,19 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
return ErrInvalidLocalIP 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 table := f.OutRules
if incoming { if incoming {
table = f.InRules table = f.InRules
} }
// We now know which firewall table to check against // We now know which firewall table to check against
if !table.match(fp, incoming, h.ConnectionState.peerCert, caPool) { if !table.match(*fp, incoming, h.ConnectionState.peerCert, caPool) {
f.metrics(incoming).droppedNoRule.Inc(1) f.metrics(incoming).droppedNoRule.Inc(1)
return ErrNoMatchingRule return ErrNoMatchingRule
} }
// We always want to conntrack since it is a faster operation // We always want to conntrack since it is a faster operation
f.addConn(fp, incoming) f.addConn(key, fp.Protocol, incoming)
return nil return nil
} }
@@ -502,9 +517,9 @@ func (f *Firewall) EmitStats() {
metrics.GetOrRegisterGauge("firewall.rules.hash", nil).Update(int64(f.GetRuleHashFNV())) metrics.GetOrRegisterGauge("firewall.rules.hash", nil).Update(int64(f.GetRuleHashFNV()))
} }
func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) bool { func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) bool {
if localCache != nil { if localCache != nil {
if _, ok := localCache[fp]; ok { if _, ok := localCache[key]; ok {
return true return true
} }
} }
@@ -517,7 +532,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
f.evict(ep) f.evict(ep)
} }
c, ok := conntrack.Conns[fp] c, ok := conntrack.Conns[key]
if !ok { if !ok {
conntrack.Unlock() conntrack.Unlock()
@@ -526,7 +541,11 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
if c.rulesVersion != f.rulesVersion { if c.rulesVersion != f.rulesVersion {
// This conntrack entry was for an older rule set, validate // This conntrack entry was for an older rule set, validate
// it still passes with the current rule set // it still passes with the current rule set. Rule matching needs
// the rich Packet form, so hydrate from key.
var fp firewall.Packet
key.Hydrate(&fp)
table := f.OutRules table := f.OutRules
if c.incoming { if c.incoming {
table = f.InRules table = f.InRules
@@ -542,7 +561,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
"oldRulesVersion", c.rulesVersion, "oldRulesVersion", c.rulesVersion,
) )
} }
delete(conntrack.Conns, fp) delete(conntrack.Conns, key)
conntrack.Unlock() conntrack.Unlock()
return false return false
} }
@@ -559,7 +578,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
c.rulesVersion = f.rulesVersion c.rulesVersion = f.rulesVersion
} }
switch fp.Protocol { switch key.Protocol {
case firewall.ProtoTCP: case firewall.ProtoTCP:
c.Expires = time.Now().Add(f.TCPTimeout) c.Expires = time.Now().Add(f.TCPTimeout)
case firewall.ProtoUDP: case firewall.ProtoUDP:
@@ -571,17 +590,17 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
conntrack.Unlock() conntrack.Unlock()
if localCache != nil { if localCache != nil {
localCache[fp] = struct{}{} localCache[key] = struct{}{}
} }
return true return true
} }
func (f *Firewall) addConn(fp firewall.Packet, incoming bool) { func (f *Firewall) addConn(key firewall.PacketKey, protocol uint8, incoming bool) {
var timeout time.Duration var timeout time.Duration
c := &conn{} c := &conn{}
switch fp.Protocol { switch protocol {
case firewall.ProtoTCP: case firewall.ProtoTCP:
timeout = f.TCPTimeout timeout = f.TCPTimeout
case firewall.ProtoUDP: case firewall.ProtoUDP:
@@ -592,9 +611,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
conntrack := f.Conntrack conntrack := f.Conntrack
conntrack.Lock() conntrack.Lock()
if _, ok := conntrack.Conns[fp]; !ok { if _, ok := conntrack.Conns[key]; !ok {
conntrack.TimerWheel.Advance(time.Now()) conntrack.TimerWheel.Advance(time.Now())
conntrack.TimerWheel.Add(fp, timeout) conntrack.TimerWheel.Add(key, timeout)
} }
// Record which rulesVersion allowed this connection, so we can retest after // Record which rulesVersion allowed this connection, so we can retest after
@@ -602,16 +621,16 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
c.incoming = incoming c.incoming = incoming
c.rulesVersion = f.rulesVersion c.rulesVersion = f.rulesVersion
c.Expires = time.Now().Add(timeout) c.Expires = time.Now().Add(timeout)
conntrack.Conns[fp] = c conntrack.Conns[key] = c
conntrack.Unlock() conntrack.Unlock()
} }
// Evict checks if a conntrack entry has expired, if so it is removed, if not it is re-added to the wheel // Evict checks if a conntrack entry has expired, if so it is removed, if not it is re-added to the wheel
// Caller must own the connMutex lock! // Caller must own the connMutex lock!
func (f *Firewall) evict(p firewall.Packet) { func (f *Firewall) evict(key firewall.PacketKey) {
// Are we still tracking this conn? // Are we still tracking this conn?
conntrack := f.Conntrack conntrack := f.Conntrack
t, ok := conntrack.Conns[p] t, ok := conntrack.Conns[key]
if !ok { if !ok {
return return
} }
@@ -621,12 +640,12 @@ func (f *Firewall) evict(p firewall.Packet) {
// Timeout is in the future, re-add the timer // Timeout is in the future, re-add the timer
if newT > 0 { if newT > 0 {
conntrack.TimerWheel.Advance(time.Now()) conntrack.TimerWheel.Advance(time.Now())
conntrack.TimerWheel.Add(p, newT) conntrack.TimerWheel.Add(key, newT)
return return
} }
// This conn is done // This conn is done
delete(conntrack.Conns, p) delete(conntrack.Conns, key)
} }
func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool { func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool {
@@ -897,7 +916,7 @@ func (flc *firewallLocalCIDR) addRule(f *Firewall, localCidr string) error {
} }
if localCidr == "" { if localCidr == "" {
if len(f.unsafeNetworks) == 0 || f.defaultLocalCIDRAny { if !f.hasUnsafeNetworks || f.defaultLocalCIDRAny {
flc.Any = true flc.Any = true
return nil return nil
} }
@@ -1055,6 +1074,7 @@ func (r *rule) sanity() error {
} }
func parsePort(s string) (int32, int32, error) { func parsePort(s string) (int32, int32, error) {
var err error
const notAPort int32 = -2 const notAPort int32 = -2
if s == "any" { if s == "any" {
return firewall.PortAny, firewall.PortAny, nil return firewall.PortAny, firewall.PortAny, nil
@@ -1063,11 +1083,11 @@ func parsePort(s string) (int32, int32, error) {
return firewall.PortFragment, firewall.PortFragment, nil return firewall.PortFragment, firewall.PortFragment, nil
} }
if !strings.Contains(s, `-`) { if !strings.Contains(s, `-`) {
rPort, err := parsePortValue("", s) rPort, err := strconv.Atoi(s)
if err != nil { if err != nil {
return notAPort, notAPort, err return notAPort, notAPort, fmt.Errorf("was not a number; `%s`", s)
} }
return rPort, rPort, nil return int32(rPort), int32(rPort), nil
} }
sPorts := strings.SplitN(s, `-`, 2) sPorts := strings.SplitN(s, `-`, 2)
@@ -1078,40 +1098,22 @@ func parsePort(s string) (int32, int32, error) {
return notAPort, notAPort, fmt.Errorf("appears to be a range but could not be parsed; `%s`", s) return notAPort, notAPort, fmt.Errorf("appears to be a range but could not be parsed; `%s`", s)
} }
startPort, err := parsePortValue("beginning range ", sPorts[0]) rStartPort, err := strconv.Atoi(sPorts[0])
if err != nil { if err != nil {
return notAPort, notAPort, err return notAPort, notAPort, fmt.Errorf("beginning range was not a number; `%s`", sPorts[0])
} }
endPort, err := parsePortValue("ending range ", sPorts[1]) rEndPort, err := strconv.Atoi(sPorts[1])
if err != nil { if err != nil {
return notAPort, notAPort, err return notAPort, notAPort, fmt.Errorf("ending range was not a number; `%s`", sPorts[1])
} }
startPort := int32(rStartPort)
endPort := int32(rEndPort)
if startPort == firewall.PortAny { if startPort == firewall.PortAny {
endPort = firewall.PortAny endPort = firewall.PortAny
} }
return startPort, endPort, nil return startPort, endPort, nil
} }
// parsePortValue accepts a base-10 decimal in [0, 65535] and returns it
// widened to int32. Using strconv.ParseUint with bitSize 16 rejects
// negative input, out-of-range input (>65535), and any non-decimal byte
// by construction, so the int32 widening that follows is provably safe
// and cannot collide with firewall.PortAny (0) or firewall.PortFragment
// (-1) via integer truncation.
//
// prefix is prepended to both error messages so callers can disambiguate
// the single-port path (prefix="") from the range bounds (prefix="beginning
// range " / "ending range "), preserving the historical error strings.
func parsePortValue(prefix, s string) (int32, error) {
n, err := strconv.ParseUint(s, 10, 16)
if err == nil {
return int32(n), nil
}
if errors.Is(err, strconv.ErrRange) {
return 0, fmt.Errorf("%sout of range [0,65535]; `%s`", prefix, s)
}
return 0, fmt.Errorf("%swas not a number; `%s`", prefix, s)
}
+4 -2
View File
@@ -10,8 +10,10 @@ import (
) )
// ConntrackCache is used as a local routine cache to know if a given flow // ConntrackCache is used as a local routine cache to know if a given flow
// has been seen in the conntrack table. // has been seen in the conntrack table. Keyed on PacketKey (dense form)
type ConntrackCache map[Packet]struct{} // rather than Packet so the lookup hashes raw bytes instead of the
// unique.Handle each netip.Addr in Packet carries.
type ConntrackCache map[PacketKey]struct{}
type ConntrackCacheTicker struct { type ConntrackCacheTicker struct {
cacheV uint64 cacheV uint64
+1 -1
View File
@@ -23,7 +23,7 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
cache: make(ConntrackCache, cacheLen), cache: make(ConntrackCache, cacheLen),
} }
for i := 0; i < cacheLen; i++ { for i := 0; i < cacheLen; i++ {
c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{} c.cache[PacketKey{LocalPort: uint16(i) + 1}] = struct{}{}
} }
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
return c return c
+74
View File
@@ -19,6 +19,25 @@ const (
PortFragment = -1 // Special value for matching `port: fragment` PortFragment = -1 // Special value for matching `port: fragment`
) )
// PacketKey is the firewall's conntrack and ConntrackCache map key — the
// dense form of the 5-tuple plus the protocol and fragment flag the
// firewall actually discriminates flows on. Kept separate from Packet so
// the conntrack-hit fast path doesn't pay for hashing the unique.Handle
// each netip.Addr carries, and so the inbound parser can skip the
// AddrFrom4/AddrFrom16 calls until rule matching actually needs them.
//
// Superset of the coalescer's flowKey shape (same 5-tuple, just in
// Local/Remote orientation rather than wire src/dst).
type PacketKey struct {
LocalAddr [16]byte
RemoteAddr [16]byte
LocalPort uint16
RemotePort uint16
IsV6 bool
Protocol uint8
Fragment bool
}
type Packet struct { type Packet struct {
LocalAddr netip.Addr LocalAddr netip.Addr
RemoteAddr netip.Addr RemoteAddr netip.Addr
@@ -31,6 +50,61 @@ type Packet struct {
Fragment bool Fragment bool
} }
// Key derives a PacketKey from a populated Packet. Used by the few code
// paths that have a Packet but no Key in hand (e.g. tests). Both inbound
// and outbound production parsers write straight into a PacketKey via
// batch.ParsePacket, so this function is rarely on the hot path.
func (fp *Packet) Key() PacketKey {
k := PacketKey{
Protocol: fp.Protocol,
Fragment: fp.Fragment,
}
k.LocalPort = fp.LocalPort
k.RemotePort = fp.RemotePort
k.IsV6 = !fp.LocalAddr.Is4()
if k.IsV6 {
k.LocalAddr = fp.LocalAddr.As16()
k.RemoteAddr = fp.RemoteAddr.As16()
} else {
v4 := fp.LocalAddr.As4()
copy(k.LocalAddr[:4], v4[:])
v4 = fp.RemoteAddr.As4()
copy(k.RemoteAddr[:4], v4[:])
}
return k
}
// Hydrate fills fp's netip.Addr fields and copies the rest from k. Called
// by the firewall slow path when conntrack misses and rule matching needs
// the rich Packet form (CIDR lookups, family checks). The fast path skips
// this entirely.
func (k *PacketKey) Hydrate(fp *Packet) {
fp.LocalPort = k.LocalPort
fp.RemotePort = k.RemotePort
fp.Protocol = k.Protocol
fp.Fragment = k.Fragment
if k.IsV6 {
fp.LocalAddr = netip.AddrFrom16(k.LocalAddr)
fp.RemoteAddr = netip.AddrFrom16(k.RemoteAddr)
} else {
var v4 [4]byte
copy(v4[:], k.LocalAddr[:4])
fp.LocalAddr = netip.AddrFrom4(v4)
copy(v4[:], k.RemoteAddr[:4])
fp.RemoteAddr = netip.AddrFrom4(v4)
}
}
func (k *PacketKey) GetRemoteAddr() netip.Addr {
if k.IsV6 {
return netip.AddrFrom16(k.RemoteAddr)
} else {
var v4 [4]byte
copy(v4[:], k.RemoteAddr[:4])
return netip.AddrFrom4(v4)
}
}
func (fp *Packet) Copy() *Packet { func (fp *Packet) Copy() *Packet {
return &Packet{ return &Packet{
LocalAddr: fp.LocalAddr, LocalAddr: fp.LocalAddr,
+53 -275
View File
@@ -211,44 +211,44 @@ func TestFirewall_Drop(t *testing.T) {
cp := cert.NewCAPool() cp := cert.NewCAPool()
// Drop outbound // Drop outbound
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil)) assert.Equal(t, ErrNoMatchingRule, fw.Drop(p.Key(), &p, false, &h, cp, nil))
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
// Allow outbound because conntrack // Allow outbound because conntrack
require.NoError(t, fw.Drop(p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
// test remote mismatch // test remote mismatch
oldRemote := p.RemoteAddr oldRemote := p.RemoteAddr
p.RemoteAddr = netip.MustParseAddr("1.2.3.10") p.RemoteAddr = netip.MustParseAddr("1.2.3.10")
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP) assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrInvalidRemoteIP)
p.RemoteAddr = oldRemote p.RemoteAddr = oldRemote
// ensure signer doesn't get in the way of group checks // ensure signer doesn't get in the way of group checks
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
// test caSha doesn't drop on match // test caSha doesn't drop on match
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
// ensure ca name doesn't get in the way of group checks // ensure ca name doesn't get in the way of group checks
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
// test caName doesn't drop on match // test caName doesn't drop on match
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
} }
func TestFirewall_DropV6(t *testing.T) { func TestFirewall_DropV6(t *testing.T) {
@@ -289,44 +289,44 @@ func TestFirewall_DropV6(t *testing.T) {
cp := cert.NewCAPool() cp := cert.NewCAPool()
// Drop outbound // Drop outbound
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil)) assert.Equal(t, ErrNoMatchingRule, fw.Drop(p.Key(), &p, false, &h, cp, nil))
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
// Allow outbound because conntrack // Allow outbound because conntrack
require.NoError(t, fw.Drop(p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
// test remote mismatch // test remote mismatch
oldRemote := p.RemoteAddr oldRemote := p.RemoteAddr
p.RemoteAddr = netip.MustParseAddr("fd12::56") p.RemoteAddr = netip.MustParseAddr("fd12::56")
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP) assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrInvalidRemoteIP)
p.RemoteAddr = oldRemote p.RemoteAddr = oldRemote
// ensure signer doesn't get in the way of group checks // ensure signer doesn't get in the way of group checks
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
// test caSha doesn't drop on match // test caSha doesn't drop on match
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
// ensure ca name doesn't get in the way of group checks // ensure ca name doesn't get in the way of group checks
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
// test caName doesn't drop on match // test caName doesn't drop on match
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}} cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
} }
func BenchmarkFirewallTable_match(b *testing.B) { func BenchmarkFirewallTable_match(b *testing.B) {
@@ -533,10 +533,10 @@ func TestFirewall_Drop2(t *testing.T) {
cp := cert.NewCAPool() cp := cert.NewCAPool()
// h1/c1 lacks the proper groups // h1/c1 lacks the proper groups
require.ErrorIs(t, fw.Drop(p, true, &h1, cp, nil), ErrNoMatchingRule) require.ErrorIs(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil), ErrNoMatchingRule)
// c has the proper groups // c has the proper groups
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
} }
func TestFirewall_Drop3(t *testing.T) { func TestFirewall_Drop3(t *testing.T) {
@@ -613,18 +613,18 @@ func TestFirewall_Drop3(t *testing.T) {
cp := cert.NewCAPool() cp := cert.NewCAPool()
// c1 should pass because host match // c1 should pass because host match
require.NoError(t, fw.Drop(p, true, &h1, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil))
// c2 should pass because ca sha match // c2 should pass because ca sha match
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(p, true, &h2, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h2, cp, nil))
// c3 should fail because no match // c3 should fail because no match
resetConntrack(fw) resetConntrack(fw)
assert.Equal(t, fw.Drop(p, true, &h3, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, true, &h3, cp, nil), ErrNoMatchingRule)
// Test a remote address match // Test a remote address match
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", ""))
require.NoError(t, fw.Drop(p, true, &h1, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil))
} }
func TestFirewall_Drop3V6(t *testing.T) { func TestFirewall_Drop3V6(t *testing.T) {
@@ -661,7 +661,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
cp := cert.NewCAPool() cp := cert.NewCAPool()
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", "")) require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", ""))
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
} }
func TestFirewall_DropConntrackReload(t *testing.T) { func TestFirewall_DropConntrackReload(t *testing.T) {
@@ -702,12 +702,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
cp := cert.NewCAPool() cp := cert.NewCAPool()
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
// Allow outbound because conntrack // Allow outbound because conntrack
require.NoError(t, fw.Drop(p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
oldFw := fw oldFw := fw
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
@@ -716,7 +716,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
fw.rulesVersion = oldFw.rulesVersion + 1 fw.rulesVersion = oldFw.rulesVersion + 1
// Allow outbound because conntrack and new rules allow port 10 // Allow outbound because conntrack and new rules allow port 10
require.NoError(t, fw.Drop(p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
oldFw = fw oldFw = fw
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate) fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
@@ -725,7 +725,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
fw.rulesVersion = oldFw.rulesVersion + 1 fw.rulesVersion = oldFw.rulesVersion + 1
// Drop outbound because conntrack doesn't match new ruleset // Drop outbound because conntrack doesn't match new ruleset
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrNoMatchingRule)
} }
func TestFirewall_ICMPPortBehavior(t *testing.T) { func TestFirewall_ICMPPortBehavior(t *testing.T) {
@@ -770,12 +770,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 0 p.LocalPort = 0
p.RemotePort = 0 p.RemotePort = 0
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
//now also allow outbound //now also allow outbound
require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
}) })
t.Run("nonzero ports", func(t *testing.T) { t.Run("nonzero ports", func(t *testing.T) {
@@ -783,12 +783,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 0xabcd p.LocalPort = 0xabcd
p.RemotePort = 0x1234 p.RemotePort = 0x1234
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
//now also allow outbound //now also allow outbound
require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
}) })
}) })
@@ -800,12 +800,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 0 p.LocalPort = 0
p.RemotePort = 0 p.RemotePort = 0
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, true, &h, cp, nil), ErrNoMatchingRule)
//now also allow outbound //now also allow outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
}) })
t.Run("nonzero ports, still blocked", func(t *testing.T) { t.Run("nonzero ports, still blocked", func(t *testing.T) {
@@ -813,12 +813,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 0xabcd p.LocalPort = 0xabcd
p.RemotePort = 0x1234 p.RemotePort = 0x1234
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, true, &h, cp, nil), ErrNoMatchingRule)
//now also allow outbound //now also allow outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
}) })
t.Run("nonzero, matching ports, still blocked", func(t *testing.T) { t.Run("nonzero, matching ports, still blocked", func(t *testing.T) {
@@ -826,12 +826,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 80 p.LocalPort = 80
p.RemotePort = 80 p.RemotePort = 80
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, true, &h, cp, nil), ErrNoMatchingRule)
//now also allow outbound //now also allow outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
}) })
}) })
t.Run("Any proto, any port", func(t *testing.T) { t.Run("Any proto, any port", func(t *testing.T) {
@@ -843,12 +843,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 0 p.LocalPort = 0
p.RemotePort = 0 p.RemotePort = 0
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
//now also allow outbound //now also allow outbound
require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
}) })
t.Run("nonzero ports, allowed", func(t *testing.T) { t.Run("nonzero ports, allowed", func(t *testing.T) {
@@ -857,15 +857,15 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
p.LocalPort = 0xabcd p.LocalPort = 0xabcd
p.RemotePort = 0x1234 p.RemotePort = 0x1234
// Drop outbound // Drop outbound
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
// Allow inbound // Allow inbound
resetConntrack(fw) resetConntrack(fw)
require.NoError(t, fw.Drop(*p, true, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
//now also allow outbound //now also allow outbound
require.NoError(t, fw.Drop(*p, false, &h, cp, nil)) require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
//different ID is blocked //different ID is blocked
p.RemotePort++ p.RemotePort++
require.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule) require.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
}) })
}) })
@@ -913,160 +913,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
Protocol: firewall.ProtoUDP, Protocol: firewall.ProtoUDP,
Fragment: false, Fragment: false,
} }
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP) assert.Equal(t, fw.Drop(p.Key(), &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) { func BenchmarkLookup(b *testing.B) {
@@ -1182,75 +1029,6 @@ func Test_parsePort(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
} }
// Test_parsePort_invalid covers inputs that must error. The named bug is
// that int32(strconv.Atoi("4294967296")) truncates to 0 == firewall.PortAny,
// silently turning a typo into a match-all-ports rule; the rest are
// representative syntax/range probes.
func Test_parsePort_invalid(t *testing.T) {
tests := []struct {
name string
input string
wantErrContains string
}{
// Numeric overflow (the named bug + boundary).
{"named bug: 2^32 truncates to PortAny", "4294967296", "out of range"},
{"just above max real port", "65536", "out of range"},
// Negatives route through the range branch and hit the empty-half
// guard; included as defense in depth so a future refactor cannot
// accidentally reach the int32 cast.
{"negative", "-1", "could not be parsed"},
// Syntax probes.
{"NUL between digits", "4\x002", "was not a number"},
{"hex notation", "0x10", "was not a number"},
{"scientific notation", "1e3", "was not a number"},
{"leading whitespace", " 42", "was not a number"},
{"fullwidth digits", "42", "was not a number"},
// Range branch.
{"range upper out of range", "1-65536", "ending range out of range"},
{"range lower out of range", "65536-65537", "beginning range out of range"},
{"range with negative upper", "1--1", "ending range was not a number"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
_, _, err := parsePort(tc.input)
require.Error(t, err, "input %q must error", tc.input)
require.ErrorContains(t, err, tc.wantErrContains)
})
}
}
// Test_parsePort_valid_boundaries locks in success cases at 0, 1, and 65535
// so a future refactor cannot regress the boundaries.
func Test_parsePort_valid_boundaries(t *testing.T) {
tests := []struct {
name string
input string
wantStart int32
wantEnd int32
}{
{"zero is PortAny", "0", 0, 0},
{"min real port", "1", 1, 1},
{"max real port", "65535", 65535, 65535},
{"range zero to max forces end to zero", "0-65535", 0, 0},
{"range max to max", "65535-65535", 65535, 65535},
{"range one to max", "1-65535", 1, 65535},
{"range with whitespace inside", " 1 - 2 ", 1, 2},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
s, e, err := parsePort(tc.input)
require.NoError(t, err)
assert.Equal(t, tc.wantStart, s, "start port")
assert.Equal(t, tc.wantEnd, e, "end port")
})
}
}
func TestNewFirewallFromConfig(t *testing.T) { func TestNewFirewallFromConfig(t *testing.T) {
l := test.NewLogger() l := test.NewLogger()
// Test a bad rule definition // Test a bad rule definition
@@ -1549,7 +1327,7 @@ func (c *testcase) Test(t *testing.T, fw *Firewall) {
t.Helper() t.Helper()
cp := cert.NewCAPool() cp := cert.NewCAPool()
resetConntrack(fw) resetConntrack(fw)
err := fw.Drop(c.p, true, c.h, cp, nil) err := fw.Drop(c.p.Key(), &c.p, true, c.h, cp, nil)
if c.err == nil { if c.err == nil {
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr) require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
} else { } else {
@@ -1741,6 +1519,6 @@ func (mf *mockFirewall) AddRule(incoming bool, proto uint8, startPort int32, end
func resetConntrack(fw *Firewall) { func resetConntrack(fw *Firewall) {
fw.Conntrack.Lock() fw.Conntrack.Lock()
fw.Conntrack.Conns = map[firewall.Packet]*conn{} fw.Conntrack.Conns = map[firewall.PacketKey]*conn{}
fw.Conntrack.Unlock() fw.Conntrack.Unlock()
} }
+11 -10
View File
@@ -9,10 +9,10 @@ require (
github.com/armon/go-radix v1.0.0 github.com/armon/go-radix v1.0.0
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432 github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
github.com/flynn/noise v1.1.0 github.com/flynn/noise v1.1.0
github.com/gaissmai/bart v0.28.0 github.com/gaissmai/bart v0.26.1
github.com/gogo/protobuf v1.3.2 github.com/gogo/protobuf v1.3.2
github.com/google/gopacket v1.1.19 github.com/google/gopacket v1.1.19
github.com/kardianos/service v1.3.0 github.com/kardianos/service v1.2.4
github.com/miekg/dns v1.1.72 github.com/miekg/dns v1.1.72
github.com/miekg/pkcs11 v1.1.2 github.com/miekg/pkcs11 v1.1.2
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
@@ -24,15 +24,15 @@ require (
github.com/vishvananda/netlink v1.3.1 github.com/vishvananda/netlink v1.3.1
go.uber.org/goleak v1.3.0 go.uber.org/goleak v1.3.0
go.yaml.in/yaml/v3 v3.0.4 go.yaml.in/yaml/v3 v3.0.4
golang.org/x/crypto v0.53.0 golang.org/x/crypto v0.50.0
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
golang.org/x/net v0.56.0 golang.org/x/net v0.53.0
golang.org/x/sync v0.21.0 golang.org/x/sync v0.20.0
golang.org/x/sys v0.46.0 golang.org/x/sys v0.43.0
golang.org/x/term v0.44.0 golang.org/x/term v0.42.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
golang.zx2c4.com/wireguard/windows v1.0.1 golang.zx2c4.com/wireguard/windows v0.6.1
google.golang.org/protobuf v1.36.11 google.golang.org/protobuf v1.36.11
gopkg.in/yaml.v3 v3.0.1 gopkg.in/yaml.v3 v3.0.1
gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe
@@ -43,6 +43,7 @@ require (
github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/davecgh/go-spew v1.1.1 // indirect github.com/davecgh/go-spew v1.1.1 // indirect
github.com/google/btree v1.1.2 // indirect github.com/google/btree v1.1.2 // indirect
github.com/guptarohit/asciigraph v0.9.0 // indirect
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/client_model v0.6.2 // indirect
@@ -50,7 +51,7 @@ require (
github.com/prometheus/procfs v0.16.1 // indirect github.com/prometheus/procfs v0.16.1 // indirect
github.com/vishvananda/netns v0.0.5 // indirect github.com/vishvananda/netns v0.0.5 // indirect
go.yaml.in/yaml/v2 v2.4.2 // indirect go.yaml.in/yaml/v2 v2.4.2 // indirect
golang.org/x/mod v0.36.0 // indirect golang.org/x/mod v0.34.0 // indirect
golang.org/x/time v0.5.0 // indirect golang.org/x/time v0.5.0 // indirect
golang.org/x/tools v0.45.0 // indirect golang.org/x/tools v0.43.0 // indirect
) )
+22 -20
View File
@@ -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/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 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag= github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
github.com/gaissmai/bart v0.28.0 h1:89yZLo8NmyqD0RYgJ3QO9HhqqGGw+oWhf90cZm69Lko= github.com/gaissmai/bart v0.26.1 h1:+w4rnLGNlA2GDVn382Tfe3jOsK5vOr5n4KmigJ9lbTo=
github.com/gaissmai/bart v0.28.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c= github.com/gaissmai/bart v0.26.1/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.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
github.com/go-kit/kit v0.9.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= github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
@@ -60,14 +60,16 @@ github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
github.com/google/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8= github.com/google/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8=
github.com/google/gopacket v1.1.19/go.mod h1:iJ8V8n6KS+z2U1A8pUwu8bW5SyEMkXJB8Yo/Vo+TKTo= github.com/google/gopacket v1.1.19/go.mod h1:iJ8V8n6KS+z2U1A8pUwu8bW5SyEMkXJB8Yo/Vo+TKTo=
github.com/guptarohit/asciigraph v0.9.0 h1:MvCSRRVkT2XvU1IO6n92o7l7zqx1DiFaoszOUZQztbY=
github.com/guptarohit/asciigraph v0.9.0/go.mod h1:dYl5wwK4gNsnFf9Zp+l06rFiDZ5YtXM6x7SRWZ3KGag=
github.com/jpillora/backoff v1.0.0/go.mod h1:J/6gKK9jxlEcS3zixgDgUAsiuZ7yrSoa/FX5e0EB2j4= github.com/jpillora/backoff v1.0.0/go.mod h1:J/6gKK9jxlEcS3zixgDgUAsiuZ7yrSoa/FX5e0EB2j4=
github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU= github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU=
github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
github.com/json-iterator/go v1.1.11/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= 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.2.0/go.mod h1:SYymIcj16QtmaHHD7aYtjjsJG7VTCxuUUipMqKk8s4w=
github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM= github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM=
github.com/kardianos/service v1.3.0 h1:/LGy+xPP2TM+GLTiCZ2di7cy0Jd/qrawlTUfqKYFdTI= github.com/kardianos/service v1.2.4 h1:XNlGtZOYNx2u91urOdg/Kfmc+gfmuIo1Dd3rEi2OgBk=
github.com/kardianos/service v1.3.0/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc= github.com/kardianos/service v1.2.4/go.mod h1:E4V9ufUuY82F7Ztlu1eN9VXWIQxg8NoLQlmFe0MtrXc=
github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= 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/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck=
github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo=
@@ -162,16 +164,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-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-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY= 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/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/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.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.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
golang.org/x/mod v0.36.0 h1:JJjpVx6myfUsUdAzZuOSTTmRE0PfZeNWzzvKrP7amb4= golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
golang.org/x/mod v0.36.0/go.mod h1:moc6ELqsWcOw5Ef3xVprK5ul/MvtVvkIXLziUOICjUQ= golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= 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-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
@@ -182,8 +184,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-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-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.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= 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-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
@@ -191,8 +193,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-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-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.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sync v0.20.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-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-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
@@ -208,11 +210,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.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.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw= golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc= golang.org/x/term v0.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y= golang.org/x/term v0.42.0/go.mod h1:Dq/D+snpsbazcBG5+F9Q1n2rXV8Ma+71xEjTRufARgY=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= 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.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
@@ -223,8 +225,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-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-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.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA=
golang.org/x/tools v0.45.0 h1:18qN3FAooORvApf5XjCXgsuayZOEtXf6JK18I3+ONa8= golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
golang.org/x/tools v0.45.0/go.mod h1:LuUGqqaXcXMEFEruIVJVm5mgDD8vww/z/SR1gQ4uE/0= golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= 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-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
@@ -233,8 +235,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/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 h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4= golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8= golang.zx2c4.com/wireguard/windows v0.6.1 h1:XMaKojH1Hs/raMrmnir4n35nTvzvWj7NmSYzHn2F4qU=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs= golang.zx2c4.com/wireguard/windows v0.6.1/go.mod h1:04aqInu5GYuTFvMuDw/rKBAF7mHrltW/3rekpfbbZDM=
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4= 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-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
-1
View File
@@ -13,7 +13,6 @@ var (
ErrUnknownSubtype = errors.New("unknown handshake subtype") ErrUnknownSubtype = errors.New("unknown handshake subtype")
ErrMissingContent = errors.New("expected handshake content but message was empty") ErrMissingContent = errors.New("expected handshake content but message was empty")
ErrUnexpectedContent = errors.New("received unexpected handshake content") 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") ErrIndexAllocation = errors.New("failed to allocate local index")
ErrNoCredential = errors.New("no handshake credential available for cert version") ErrNoCredential = errors.New("no handshake credential available for cert version")
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key") ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
+2 -10
View File
@@ -312,19 +312,11 @@ func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
// Process payload // Process payload
if flags.expectsPayload { if flags.expectsPayload {
var remoteIndex uint32
if m.result.Initiator { if m.result.Initiator {
remoteIndex = payload.ResponderIndex m.result.RemoteIndex = payload.ResponderIndex
} else { } else {
remoteIndex = payload.InitiatorIndex m.result.RemoteIndex = payload.InitiatorIndex
} }
// 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.result.HandshakeTime = payload.Time
m.payloadSet = true m.payloadSet = true
} }
-18
View File
@@ -229,24 +229,6 @@ func TestMachineProcessPayload(t *testing.T) {
require.ErrorIs(t, err, ErrUnexpectedContent) require.ErrorIs(t, err, ErrUnexpectedContent)
assert.True(t, m.Failed()) 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 // TestMachineRequireComplete checks the fail-on-incomplete-handshake path
+11 -4
View File
@@ -83,7 +83,6 @@ type HandshakeHostInfo struct {
initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake? initiatingVersionOverride cert.Version // Should we use a non-default cert version for this handshake?
counter int64 // How many attempts have we made so far counter int64 // How many attempts have we made so far
lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt lastRemotes []netip.AddrPort // Remotes that we sent to during the previous attempt
lastRelays []netip.Addr // Relays we attempted to use during the previous attempt
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
hostinfo *HostInfo hostinfo *HostInfo
@@ -218,6 +217,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
fields := []any{ fields := []any{
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()), "udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
"initiatorIndex", hh.hostinfo.localIndexId, "initiatorIndex", hh.hostinfo.localIndexId,
"remoteIndex", hh.hostinfo.remoteIndexId,
"durationNs", time.Since(hh.startTime).Nanoseconds(), "durationNs", time.Since(hh.startTime).Nanoseconds(),
} }
// hh.machine can be nil here if buildStage0Packet never succeeded // hh.machine can be nil here if buildStage0Packet never succeeded
@@ -323,7 +323,7 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
) )
} }
hm.f.relayManager.StartRelays(hm.f, vpnIp, hh, stage0) hm.f.relayManager.StartRelays(hm.f, vpnIp, hostinfo, stage0)
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add // If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
if !lighthouseTriggered { if !lighthouseTriggered {
@@ -430,11 +430,14 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
// Check if we already have a tunnel with this vpn ip // Check if we already have a tunnel with this vpn ip
existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]] existingHostInfo, found := hm.mainHostMap.Hosts[hostinfo.vpnAddrs[0]]
if found && existingHostInfo != nil { if found && existingHostInfo != nil {
// Is it just a delayed handshake packet? Check every hostinfo we hold for this address. testHostInfo := existingHostInfo
for _, testHostInfo := range hm.mainHostMap.unlockedGetHostList(hostinfo.vpnAddrs[0]) { for testHostInfo != nil {
// Is it just a delayed handshake packet?
if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) { if bytes.Equal(hostinfo.HandshakePacket[handshakePacket], testHostInfo.HandshakePacket[handshakePacket]) {
return testHostInfo, ErrAlreadySeen return testHostInfo, ErrAlreadySeen
} }
testHostInfo = testHostInfo.next
} }
// Is this a newer handshake? // Is this a newer handshake?
@@ -462,6 +465,7 @@ func (hm *HandshakeManager) CheckAndComplete(hostinfo *HostInfo, handshakePacket
// We have a collision, but this can happen since we can't control // We have a collision, but this can happen since we can't control
// the remote ID. Just log about the situation as a note. // the remote ID. Just log about the situation as a note.
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex", hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
"remoteIndex", hostinfo.remoteIndexId,
"collision", existingRemoteIndex.vpnAddrs, "collision", existingRemoteIndex.vpnAddrs,
) )
} }
@@ -484,6 +488,7 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
// We have a collision, but this can happen since we can't control // We have a collision, but this can happen since we can't control
// the remote ID. Just log about the situation as a note. // the remote ID. Just log about the situation as a note.
hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex", hostinfo.logger(hm.l).Info("New host shadows existing host remoteIndex",
"remoteIndex", hostinfo.remoteIndexId,
"collision", existingRemoteIndex.vpnAddrs, "collision", existingRemoteIndex.vpnAddrs,
) )
} }
@@ -793,6 +798,7 @@ func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *head
} }
hm.sendHandshakeResponse(via, response, hostinfo, false) hm.sendHandshakeResponse(via, response, hostinfo, false)
f.connectionManager.AddTrafficWatch(hostinfo)
hostinfo.remotes.RefreshFromHandshake(vpnAddrs) hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
// Don't wait for UpdateWorker // Don't wait for UpdateWorker
@@ -959,6 +965,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate) hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
hm.Complete(hostinfo, f) hm.Complete(hostinfo, f)
f.connectionManager.AddTrafficWatch(hostinfo)
if len(hh.packetStore) > 0 { if len(hh.packetStore) > 0 {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
+106 -159
View File
@@ -60,16 +60,7 @@ type HostMap struct {
Indexes map[uint32]*HostInfo Indexes map[uint32]*HostInfo
Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object Relays map[uint32]*HostInfo // Maps a Relay IDX to a Relay HostInfo object
RemoteIndexes map[uint32]*HostInfo RemoteIndexes map[uint32]*HostInfo
// Hosts maps a vpn address to its primary hostinfo, one entry per address we hold a tunnel
// for. moreHosts only has an entry while an address is held by 2 or more hostinfos and stores
// the full most-recent-first list; moreHosts[a][0] is always the same hostinfo as Hosts[a].
// Each address gets its own independent list, so a hostinfo owning multiple addresses can
// never corrupt another address's ordering the way the old shared next/prev chain could.
// Entries in moreHosts are only ever written by unlockedSetHostsForAddr; Hosts is written
// directly only in the single-hostinfo fast paths where moreHosts is known to have no entry,
// and unlockedDeleteHostInfo swaps either map for a fresh one when it fully drains.
Hosts map[netip.Addr]*HostInfo Hosts map[netip.Addr]*HostInfo
moreHosts map[netip.Addr][]*HostInfo
preferredRanges atomic.Pointer[[]netip.Prefix] preferredRanges atomic.Pointer[[]netip.Prefix]
l *slog.Logger l *slog.Logger
} }
@@ -147,9 +138,9 @@ func (rs *RelayState) InsertRelayTo(ip netip.Addr) {
} }
func (rs *RelayState) CopyRelayIps() []netip.Addr { func (rs *RelayState) CopyRelayIps() []netip.Addr {
ret := make([]netip.Addr, len(rs.relays))
rs.RLock() rs.RLock()
defer rs.RUnlock() defer rs.RUnlock()
ret := make([]netip.Addr, len(rs.relays))
copy(ret, rs.relays) copy(ret, rs.relays)
return ret return ret
} }
@@ -238,7 +229,7 @@ const (
) )
type HostInfo struct { type HostInfo struct {
remote atomic.Pointer[netip.AddrPort] remote netip.AddrPort
remotes *RemoteList remotes *RemoteList
promoteCounter atomic.Uint32 promoteCounter atomic.Uint32
ConnectionState *ConnectionState ConnectionState *ConnectionState
@@ -275,6 +266,10 @@ type HostInfo struct {
lastRoam time.Time lastRoam time.Time
lastRoamRemote netip.AddrPort lastRoamRemote netip.AddrPort
// Used to track other hostinfos for this vpn ip since only 1 can be primary
// Synchronised via hostmap lock and not the hostinfo lock.
next, prev *HostInfo
//TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing //TODO: in, out, and others might benefit from being an atomic.Int32. We could collapse connectionManager pendingDeletion, relayUsed, and in/out into this 1 thing
in, out, pendingDeletion atomic.Bool in, out, pendingDeletion atomic.Bool
@@ -339,7 +334,6 @@ func newHostMap(l *slog.Logger) *HostMap {
Relays: map[uint32]*HostInfo{}, Relays: map[uint32]*HostInfo{},
RemoteIndexes: map[uint32]*HostInfo{}, RemoteIndexes: map[uint32]*HostInfo{},
Hosts: map[netip.Addr]*HostInfo{}, Hosts: map[netip.Addr]*HostInfo{},
moreHosts: map[netip.Addr][]*HostInfo{},
l: l, l: l,
} }
} }
@@ -388,55 +382,13 @@ func (hm *HostMap) EmitStats() {
metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen)) metrics.GetOrRegisterGauge("hostmap.main.relayIndexes", nil).Update(int64(relaysLen))
} }
// unlockedSetHostsForAddr stores the per-address hostinfo list (list[0] is the primary). An empty // DeleteHostInfo will fully unlink the hostinfo and return true if it was the final hostinfo for this vpn ip
// list removes the address. This is the one place Hosts and moreHosts are written together, keep
// it that way. Callers must hold the write lock.
func (hm *HostMap) unlockedSetHostsForAddr(addr netip.Addr, list []*HostInfo) {
if len(list) == 0 {
delete(hm.Hosts, addr)
delete(hm.moreHosts, addr)
return
}
hm.Hosts[addr] = list[0]
if len(list) > 1 {
hm.moreHosts[addr] = list
} else {
delete(hm.moreHosts, addr)
}
}
// unlockedGetHostList returns every hostinfo holding addr, primary first, or nil if we have no
// tunnel for addr. The common single-hostinfo case builds a fresh one element list, so keep this
// off the packet hot path; the primary is a direct Hosts read. Callers must hold the lock (read
// or write).
func (hm *HostMap) unlockedGetHostList(addr netip.Addr) []*HostInfo {
if list, ok := hm.moreHosts[addr]; ok {
return list
}
if h, ok := hm.Hosts[addr]; ok {
return []*HostInfo{h}
}
return nil
}
// removeHostInfo returns list with hi removed (order preserved), or list unchanged if hi is
// absent. It deletes in place: every mutator holds the hostmap write lock and no reader ever
// retains a slice across a mutation (readers iterate under RLock), so there is no snapshot to
// invalidate.
func removeHostInfo(list []*HostInfo, hi *HostInfo) []*HostInfo {
idx := slices.Index(list, hi)
if idx < 0 {
return list
}
return slices.Delete(list, idx, idx+1)
}
// DeleteHostInfo will fully unlink the hostinfo and return true if no other hostinfo still holds
// any of its vpn addrs, meaning we no longer have a tunnel to the peer
func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool { func (hm *HostMap) DeleteHostInfo(hostinfo *HostInfo) bool {
// Delete the host itself, ensuring it's not modified anymore // Delete the host itself, ensuring it's not modified anymore
hm.Lock() hm.Lock()
final := hm.unlockedDeleteHostInfo(hostinfo) // If we have a previous or next hostinfo then we are not the last one for this vpn ip
final := (hostinfo.next == nil && hostinfo.prev == nil)
hm.unlockedDeleteHostInfo(hostinfo)
hm.Unlock() hm.Unlock()
return final return final
@@ -448,67 +400,86 @@ func (hm *HostMap) MakePrimary(hostinfo *HostInfo) {
hm.unlockedMakePrimary(hostinfo) hm.unlockedMakePrimary(hostinfo)
} }
// unlockedMakePrimary reports whether hostinfo is (now) the primary for each of its addresses, func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) {
// false only when it is no longer in the hostmap at all. // Get the current primary, if it exists
func (hm *HostMap) unlockedMakePrimary(hostinfo *HostInfo) bool { oldHostinfo := hm.Hosts[hostinfo.vpnAddrs[0]]
// A hostinfo that is no longer in the hostmap must not be re-inserted here. Callers can race
// tunnel teardown, deciding to promote under the read lock and only taking the write lock // Every address in the hostinfo gets elevated to primary
// after a delete fully unlinked the hostinfo (connection manager swapPrimary, AddRelay). Every for _, vpnAddr := range hostinfo.vpnAddrs {
// live hostinfo is registered in Indexes by unlockedAddHostInfo, so this is a membership test. //NOTE: It is possible that we leave a dangling hostinfo here but connection manager works on
if hm.Indexes[hostinfo.localIndexId] != hostinfo { // indexes so it should be fine.
return false hm.Hosts[vpnAddr] = hostinfo
} }
// Move hostinfo to the front (primary) of each of its address lists. The lists are // If we are already primary then we won't bother re-linking
// independent per address, so this can never leave a dangling entry the way promoting if oldHostinfo == hostinfo {
// against a single shared chain could. return
for _, addr := range hostinfo.vpnAddrs {
if hm.Hosts[addr] == hostinfo {
// Already primary for this address, the list is already in the right order
continue
}
list := removeHostInfo(hm.unlockedGetHostList(addr), hostinfo)
list = append([]*HostInfo{hostinfo}, list...)
hm.unlockedSetHostsForAddr(addr, list)
}
return true
} }
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index // Unlink this hostinfo
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have if hostinfo.prev != nil {
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse hostinfo.prev.next = hostinfo.next
// state and disestablish relays.
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
// Remove this hostinfo from each of its address lists. The lists are independent, so a
// sibling is never promoted to an address it does not own and no other list is touched.
final := true
for _, addr := range hostinfo.vpnAddrs {
if list, ok := hm.moreHosts[addr]; ok {
list = removeHostInfo(list, hostinfo)
hm.unlockedSetHostsForAddr(addr, list)
if len(list) > 0 {
final = false
} }
} else if existing, ok := hm.Hosts[addr]; ok { if hostinfo.next != nil {
if existing == hostinfo { hostinfo.next.prev = hostinfo.prev
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up. }
// If there wasn't a previous primary then clear out any links
if oldHostinfo == nil {
hostinfo.next = nil
hostinfo.prev = nil
return
}
// Relink the hostinfo as primary
hostinfo.next = oldHostinfo
oldHostinfo.prev = hostinfo
hostinfo.prev = nil
}
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
for _, addr := range hostinfo.vpnAddrs {
h := hm.Hosts[addr]
for h != nil {
if h == hostinfo {
hm.unlockedInnerDeleteHostInfo(h, addr)
}
h = h.next
}
}
}
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) delete(hm.Hosts, addr)
} else {
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
final = false
}
}
}
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
// through a large peer count gives the memory back. Same idiom as the index maps below.
if len(hm.Hosts) == 0 { if len(hm.Hosts) == 0 {
hm.Hosts = map[netip.Addr]*HostInfo{} hm.Hosts = map[netip.Addr]*HostInfo{}
} }
if len(hm.moreHosts) == 0 {
hm.moreHosts = 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
}
}
hostinfo.next = nil
hostinfo.prev = nil
// The remote index uses index ids outside our control so lets make sure we are only removing // The remote index uses index ids outside our control so lets make sure we are only removing
// the remote index pointer here if it points to the hostinfo we are deleting // the remote index pointer here if it points to the hostinfo we are deleting
hostinfo2, ok := hm.RemoteIndexes[hostinfo.remoteIndexId] hostinfo2, ok := hm.RemoteIndexes[hostinfo.remoteIndexId]
@@ -531,7 +502,7 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
) )
} }
if final { if isLastHostinfo {
// I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next // I have lost connectivity to my peers. My relay tunnel is likely broken. Mark the next
// hops as 'Requested' so that new relay tunnels are created in the future. // hops as 'Requested' so that new relay tunnels are created in the future.
hm.unlockedDisestablishVpnAddrRelayFor(hostinfo) hm.unlockedDisestablishVpnAddrRelayFor(hostinfo)
@@ -540,8 +511,6 @@ func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() { for _, localRelayIdx := range hostinfo.relayState.CopyRelayForIdxs() {
delete(hm.Relays, localRelayIdx) delete(hm.Relays, localRelayIdx)
} }
return final
} }
func (hm *HostMap) QueryIndex(index uint32) *HostInfo { func (hm *HostMap) QueryIndex(index uint32) *HostInfo {
@@ -585,30 +554,19 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
hm.RLock() hm.RLock()
defer hm.RUnlock() defer hm.RUnlock()
// This runs per relayed packet, so check the primary with a single map probe and only consult
// moreHosts when the primary can't relay for us.
h, ok := hm.Hosts[relayHostIp] h, ok := hm.Hosts[relayHostIp]
if !ok { if !ok {
return nil, nil, errors.New("unable to find host") return nil, nil, errors.New("unable to find host")
} }
for h != nil {
for _, targetIp := range targetIps { for _, targetIp := range targetIps {
r, ok := h.relayState.QueryRelayForByIp(targetIp) r, ok := h.relayState.QueryRelayForByIp(targetIp)
if ok && r.State == Established { if ok && r.State == Established {
return h, r, nil return h, r, nil
} }
} }
h = h.next
if list, ok := hm.moreHosts[relayHostIp]; ok {
// list[0] is the primary we already checked
for _, h := range list[1:] {
for _, targetIp := range targetIps {
r, ok := h.relayState.QueryRelayForByIp(targetIp)
if ok && r.State == Established {
return h, r, nil
}
}
}
} }
return nil, nil, errors.New("unable to find host with relay") return nil, nil, errors.New("unable to find host with relay")
@@ -616,14 +574,20 @@ func (hm *HostMap) QueryVpnAddrsRelayFor(targetIps []netip.Addr, relayHostIp net
func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) { func (hm *HostMap) unlockedDisestablishVpnAddrRelayFor(hi *HostInfo) {
for _, relayHostIp := range hi.relayState.CopyRelayIps() { for _, relayHostIp := range hi.relayState.CopyRelayIps() {
for _, h := range hm.unlockedGetHostList(relayHostIp) { if h, ok := hm.Hosts[relayHostIp]; ok {
for h != nil {
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished) h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
h = h.next
}
} }
} }
for _, rs := range hi.relayState.CopyAllRelayFor() { for _, rs := range hi.relayState.CopyAllRelayFor() {
if rs.Type == ForwardingType { if rs.Type == ForwardingType {
for _, h := range hm.unlockedGetHostList(rs.PeerAddr) { if h, ok := hm.Hosts[rs.PeerAddr]; ok {
for h != nil {
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished) h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
h = h.next
}
} }
} }
} }
@@ -659,11 +623,6 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
hm.Indexes[hostinfo.localIndexId] = hostinfo hm.Indexes[hostinfo.localIndexId] = hostinfo
hm.RemoteIndexes[hostinfo.remoteIndexId] = 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) { if hm.l.Enabled(context.Background(), slog.LevelDebug) {
hm.l.Debug("Hostmap vpnIp added", hm.l.Debug("Hostmap vpnIp added",
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts), "hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
@@ -673,27 +632,22 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
} }
func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) { func (hm *HostMap) unlockedInnerAddHostInfo(vpnAddr netip.Addr, hostinfo *HostInfo, f *Interface) {
existing, ok := hm.Hosts[vpnAddr] existing := hm.Hosts[vpnAddr]
if !ok {
// Common case, the first hostinfo for this address. moreHosts stays empty.
hm.Hosts[vpnAddr] = hostinfo hm.Hosts[vpnAddr] = hostinfo
return
if existing != nil && existing != hostinfo {
hostinfo.next = existing
existing.prev = hostinfo
} }
// The new hostinfo becomes the primary for this address. Remove any stale copy of it first so i := 1
// we never hold a duplicate, then prepend. check := hostinfo
list, ok := hm.moreHosts[vpnAddr] for check != nil {
if !ok { if i > MaxHostInfosPerVpnIp {
list = []*HostInfo{existing} hm.unlockedDeleteHostInfo(check)
} }
list = removeHostInfo(list, hostinfo) check = check.next
list = append([]*HostInfo{hostinfo}, list...) i++
hm.unlockedSetHostsForAddr(vpnAddr, list)
// Enforce the per-address cap by fully retiring the oldest hostinfo once we exceed it.
// Deleting it removes it from all of its addresses and the index maps, matching prior behavior.
if len(list) > MaxHostInfosPerVpnIp {
hm.unlockedDeleteHostInfo(list[len(list)-1])
} }
} }
@@ -725,7 +679,7 @@ func (hm *HostMap) ForEachIndex(f controlEach) {
func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interface) { func (i *HostInfo) TryPromoteBest(preferredRanges []netip.Prefix, ifce *Interface) {
c := i.promoteCounter.Add(1) c := i.promoteCounter.Add(1)
if c%ifce.tryPromoteEvery.Load() == 0 { if c%ifce.tryPromoteEvery.Load() == 0 {
remote := i.GetRemote() remote := i.remote
// return early if we are already on a preferred remote // return early if we are already on a preferred remote
if remote.IsValid() { if remote.IsValid() {
@@ -767,18 +721,11 @@ func (i *HostInfo) GetCert() *cert.CachedCertificate {
return nil 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? // TODO: Maybe use ViaSender here?
func (i *HostInfo) SetRemote(remote netip.AddrPort) { func (i *HostInfo) SetRemote(remote netip.AddrPort) {
// We copy here because we likely got this remote from a source that reuses the object // We copy here because we likely got this remote from a source that reuses the object
if i.GetRemote() != remote { if i.remote != remote {
i.remote.Store(&remote) i.remote = remote
i.remotes.LearnRemote(i.vpnAddrs[0], remote) i.remotes.LearnRemote(i.vpnAddrs[0], remote)
} }
} }
@@ -790,7 +737,7 @@ func (i *HostInfo) SetRemoteIfPreferred(hm *HostMap, via ViaSender) bool {
return false return false
} }
currentRemote := i.GetRemote() currentRemote := i.remote
if !currentRemote.IsValid() { if !currentRemote.IsValid() {
i.SetRemote(via.UdpAddr) i.SetRemote(via.UdpAddr)
return true return true
+138 -295
View File
@@ -2,7 +2,6 @@ package nebula
import ( import (
"net/netip" "net/netip"
"slices"
"testing" "testing"
"github.com/slackhq/nebula/config" "github.com/slackhq/nebula/config"
@@ -11,84 +10,78 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
// chainIds returns the localIndexIds of the hostinfos holding addr, primary (index 0) first. It
// also validates the Hosts/moreHosts sync contract on every call so a mutation that broke it
// fails fast.
func chainIds(t *testing.T, hm *HostMap, addr netip.Addr) []uint32 {
t.Helper()
assertHostMapInvariants(t, hm)
list := hm.unlockedGetHostList(addr)
ids := make([]uint32, len(list))
for i, h := range list {
ids[i] = h.localIndexId
}
return ids
}
// assertHostMapInvariants checks the Hosts/moreHosts contract: moreHosts only holds addresses
// with 2 or more hostinfos, its first entry is always the primary in Hosts, lists never hold
// duplicates, every hostinfo in a list owns the address and is registered in Indexes, and every
// indexed hostinfo is reachable through each of its addresses.
func assertHostMapInvariants(t *testing.T, hm *HostMap) {
t.Helper()
for addr, list := range hm.moreHosts {
require.GreaterOrEqualf(t, len(list), 2, "moreHosts[%s] must hold at least 2 hostinfos", addr)
require.Samef(t, hm.Hosts[addr], list[0], "moreHosts[%s][0] must match the primary in Hosts", addr)
seen := map[*HostInfo]bool{}
for _, h := range list {
require.NotNilf(t, h, "moreHosts[%s] must never hold a nil hostinfo", addr)
require.Falsef(t, seen[h], "moreHosts[%s] holds hostinfo %d twice", addr, h.localIndexId)
seen[h] = true
require.Samef(t, hm.Indexes[h.localIndexId], h, "moreHosts[%s] member %d is not registered in Indexes", addr, h.localIndexId)
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "moreHosts[%s] member %d does not own the address", addr, h.localIndexId)
}
}
for addr, h := range hm.Hosts {
require.NotNilf(t, h, "Hosts[%s] must never be nil", addr)
require.Samef(t, hm.Indexes[h.localIndexId], h, "Hosts[%s] primary %d is not registered in Indexes", addr, h.localIndexId)
require.Truef(t, slices.Contains(h.vpnAddrs, addr), "Hosts[%s] primary (index %d) does not own the address", addr, h.localIndexId)
}
for idx, h := range hm.Indexes {
require.Equalf(t, idx, h.localIndexId, "Indexes[%d] holds hostinfo with localIndexId %d", idx, h.localIndexId)
for _, va := range h.vpnAddrs {
require.Truef(t, slices.Contains(hm.unlockedGetHostList(va), h), "indexed hostinfo %d is missing from the list for %s", idx, va)
}
}
}
func TestHostMap_MakePrimary(t *testing.T) { func TestHostMap_MakePrimary(t *testing.T) {
l := test.NewLogger() l := test.NewLogger()
hm := newHostMap(l) hm := newHostMap(l)
f := &Interface{} f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2} h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3} h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4} h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
hm.unlockedAddHostInfo(h4, f) hm.unlockedAddHostInfo(h4, f)
hm.unlockedAddHostInfo(h3, f) hm.unlockedAddHostInfo(h3, f)
hm.unlockedAddHostInfo(h2, f) hm.unlockedAddHostInfo(h2, f)
hm.unlockedAddHostInfo(h1, f) hm.unlockedAddHostInfo(h1, f)
// Most-recently-added is primary: h1, h2, h3, h4 // Make sure we go h1 -> h2 -> h3 -> h4
assert.Equal(t, []uint32{1, 2, 3, 4}, chainIds(t, hm, a)) prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h1, hm.QueryVpnAddr(a)) assert.Equal(t, h1.localIndexId, prim.localIndexId)
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
assert.Nil(t, h4.next)
// Swap the middle to primary: h3, h1, h2, h4 // Swap h3/middle to primary
hm.MakePrimary(h3) hm.MakePrimary(h3)
assert.Equal(t, []uint32{3, 1, 2, 4}, chainIds(t, hm, a))
assert.Equal(t, h3, hm.QueryVpnAddr(a))
// Swap the tail to primary: h4, h3, h1, h2 // Make sure we go h3 -> h1 -> h2 -> h4
hm.MakePrimary(h4) prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a)) assert.Equal(t, h3.localIndexId, prim.localIndexId)
assert.Equal(t, h1.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
assert.Nil(t, h4.next)
// Swapping the current primary again is a no-op // Swap h4/tail to primary
hm.MakePrimary(h4) hm.MakePrimary(h4)
assert.Equal(t, []uint32{4, 3, 1, 2}, chainIds(t, hm, a))
// Make sure we go h4 -> h3 -> h1 -> h2
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h4.localIndexId, prim.localIndexId)
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Nil(t, h2.next)
// Swap h4 again should be no-op
hm.MakePrimary(h4)
// Make sure we go h4 -> h3 -> h1 -> h2
prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, h4.localIndexId, prim.localIndexId)
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h3.next.localIndexId)
assert.Equal(t, h4.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h2.localIndexId, h1.next.localIndexId)
assert.Equal(t, h3.localIndexId, h1.prev.localIndexId)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Nil(t, h2.next)
} }
func TestHostMap_DeleteHostInfo(t *testing.T) { func TestHostMap_DeleteHostInfo(t *testing.T) {
@@ -96,14 +89,13 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
hm := newHostMap(l) hm := newHostMap(l)
f := &Interface{} f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} h1 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 1}
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2} h2 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 2}
h3 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 3} h3 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 3}
h4 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 4} h4 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 4}
h5 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 5} h5 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 5}
h6 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 6} h6 := &HostInfo{vpnAddrs: []netip.Addr{netip.MustParseAddr("0.0.0.1")}, localIndexId: 6}
hm.unlockedAddHostInfo(h6, f) hm.unlockedAddHostInfo(h6, f)
hm.unlockedAddHostInfo(h5, f) hm.unlockedAddHostInfo(h5, f)
@@ -112,243 +104,94 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
hm.unlockedAddHostInfo(h2, f) hm.unlockedAddHostInfo(h2, f)
hm.unlockedAddHostInfo(h1, f) hm.unlockedAddHostInfo(h1, f)
// h6 is evicted by the MaxHostInfosPerVpnIp cap; the rest are newest-first. // h6 should be deleted
assert.Nil(t, hm.QueryIndex(h6.localIndexId)) assert.Nil(t, h6.next)
assert.Equal(t, []uint32{1, 2, 3, 4, 5}, chainIds(t, hm, a)) assert.Nil(t, h6.prev)
h := hm.QueryIndex(h6.localIndexId)
assert.Nil(t, h)
// Delete primary; not final since siblings remain. // Make sure we go h1 -> h2 -> h3 -> h4 -> h5
assert.False(t, hm.DeleteHostInfo(h1)) prim := hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a)) assert.Equal(t, h1.localIndexId, prim.localIndexId)
assert.Equal(t, h2.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h1.localIndexId, h2.prev.localIndexId)
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
assert.Nil(t, h5.next)
// Deleting the same hostinfo again must not report final while siblings remain and must not // Delete primary
// disturb the list. The old chain code got this wrong: the first delete nil'd next/prev, so a hm.DeleteHostInfo(h1)
// second delete looked final and wiped lighthouse state out from under the live sibling. assert.Nil(t, h1.prev)
assert.False(t, hm.DeleteHostInfo(h1)) assert.Nil(t, h1.next)
assert.Equal(t, []uint32{2, 3, 4, 5}, chainIds(t, hm, a))
// Delete a middle node. // Make sure we go h2 -> h3 -> h4 -> h5
assert.False(t, hm.DeleteHostInfo(h3)) prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, []uint32{2, 4, 5}, chainIds(t, hm, a)) assert.Equal(t, h2.localIndexId, prim.localIndexId)
assert.Equal(t, h3.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h3.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h3.prev.localIndexId)
assert.Equal(t, h4.localIndexId, h3.next.localIndexId)
assert.Equal(t, h3.localIndexId, h4.prev.localIndexId)
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
assert.Nil(t, h5.next)
// Delete the tail. // Delete in the middle
assert.False(t, hm.DeleteHostInfo(h5)) hm.DeleteHostInfo(h3)
assert.Equal(t, []uint32{2, 4}, chainIds(t, hm, a)) assert.Nil(t, h3.prev)
assert.Nil(t, h3.next)
// Delete the head; h4 remains and becomes primary. // Make sure we go h2 -> h4 -> h5
assert.False(t, hm.DeleteHostInfo(h2)) prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, []uint32{4}, chainIds(t, hm, a)) assert.Equal(t, h2.localIndexId, prim.localIndexId)
assert.Equal(t, h4, hm.QueryVpnAddr(a)) assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
assert.Equal(t, h5.localIndexId, h4.next.localIndexId)
assert.Equal(t, h4.localIndexId, h5.prev.localIndexId)
assert.Nil(t, h5.next)
// Delete the only remaining item; final is true and the address is gone. // Delete the tail
assert.True(t, hm.DeleteHostInfo(h4)) hm.DeleteHostInfo(h5)
assert.Empty(t, chainIds(t, hm, a)) assert.Nil(t, h5.prev)
assert.Nil(t, hm.QueryVpnAddr(a)) assert.Nil(t, h5.next)
// Deleting an already-gone hostinfo is still final; nothing holds the address anymore. // Make sure we go h2 -> h4
assert.True(t, hm.DeleteHostInfo(h4)) prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Empty(t, chainIds(t, hm, a)) assert.Equal(t, h2.localIndexId, prim.localIndexId)
} assert.Equal(t, h4.localIndexId, prim.next.localIndexId)
assert.Nil(t, prim.prev)
assert.Equal(t, h4.localIndexId, h2.next.localIndexId)
assert.Equal(t, h2.localIndexId, h4.prev.localIndexId)
assert.Nil(t, h4.next)
// TestHostMap_MakePrimary_DeletedHostInfo covers promoting a hostinfo that lost a race with // Delete the head
// tunnel teardown: swapPrimary and AddRelay decide to promote while holding a stale pointer and hm.DeleteHostInfo(h2)
// only take the write lock after a delete fully unlinked the hostinfo. MakePrimary must be a assert.Nil(t, h2.prev)
// no-op, not a resurrection that installs an unmanaged primary. assert.Nil(t, h2.next)
func TestHostMap_MakePrimary_DeletedHostInfo(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
h1 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1} // Make sure we only have h4
h2 := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 2} prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
hm.unlockedAddHostInfo(h1, f) assert.Equal(t, h4.localIndexId, prim.localIndexId)
hm.unlockedAddHostInfo(h2, f) assert.Nil(t, prim.prev)
assert.Nil(t, prim.next)
assert.Nil(t, h4.next)
// h1 is fully deleted while another goroutine still holds a pointer to it. // Delete the only item
assert.False(t, hm.DeleteHostInfo(h1)) hm.DeleteHostInfo(h4)
assert.Equal(t, []uint32{2}, chainIds(t, hm, a)) assert.Nil(t, h4.prev)
assert.Nil(t, h4.next)
// The stale promote must not bring it back. // Make sure we have nil
hm.MakePrimary(h1) prim = hm.QueryVpnAddr(netip.MustParseAddr("0.0.0.1"))
assert.Equal(t, []uint32{2}, chainIds(t, hm, a)) assert.Nil(t, prim)
assert.Equal(t, h2, hm.QueryVpnAddr(a))
assert.Nil(t, hm.QueryIndex(h1.localIndexId))
}
// TestHostMap_QueryVpnAddrsRelayFor_NonPrimary makes sure a relay established on an older
// hostinfo is still found after a newer tunnel without relay state takes primary for the same
// address. The lookup checks the primary first and falls back to the rest of the list.
func TestHostMap_QueryVpnAddrsRelayFor_NonPrimary(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
relayAddr := netip.MustParseAddr("0.0.0.9")
target := netip.MustParseAddr("0.0.0.1")
older := &HostInfo{
vpnAddrs: []netip.Addr{relayAddr},
localIndexId: 1,
relayState: RelayState{
relayForByAddr: map[netip.Addr]*Relay{},
relayForByIdx: map[uint32]*Relay{},
},
}
older.relayState.InsertRelay(target, 100, &Relay{Type: ForwardingType, State: Established, LocalIndex: 100, PeerAddr: target})
hm.unlockedAddHostInfo(older, f)
// The relay is found on the primary.
h, r, err := hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
require.NoError(t, err)
assert.Equal(t, older, h)
assert.Equal(t, uint32(100), r.LocalIndex)
// A re-handshake with no relay state takes primary; the established relay on the older
// hostinfo must still be found through the fallback.
newer := &HostInfo{vpnAddrs: []netip.Addr{relayAddr}, localIndexId: 2}
hm.unlockedAddHostInfo(newer, f)
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, relayAddr))
h, r, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, relayAddr)
require.NoError(t, err)
assert.Equal(t, older, h)
assert.Equal(t, uint32(100), r.LocalIndex)
// No hostinfo at all is a plain miss.
_, _, err = hm.QueryVpnAddrsRelayFor([]netip.Addr{target}, netip.MustParseAddr("0.0.0.42"))
require.Error(t, err)
}
// TestHostMap_DeleteHostInfo_MultipleVpnAddrs exercises the case where a hostinfo carries more than one
// vpnAddr and shares its next/prev chain with a live sibling. Deleting the head must not corrupt the
// sibling: every address the sibling owns has to keep pointing at it. The pre-fix code unlinked the shared
// chain once per vpnAddr, so on the first address it nil'd next/prev, and on the second address the node
// looked already-detached: it dropped the map entry instead of promoting the sibling (and tripped the
// isLastHostinfo relay teardown). See unlockedDeleteHostInfo.
func TestHostMap_DeleteHostInfo_MultipleVpnAddrs(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
b := netip.MustParseAddr("0.0.0.2")
// Two tunnels for the same peer, each reachable at both a and b.
other := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 1}
head := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
hm.unlockedAddHostInfo(other, f)
hm.unlockedAddHostInfo(head, f)
// head is primary for both addresses, other is next in each address's list.
assert.Equal(t, head, hm.QueryVpnAddr(a))
assert.Equal(t, head, hm.QueryVpnAddr(b))
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, b))
// Delete the head. other is still live, so it must become primary for BOTH addresses.
assert.False(t, hm.DeleteHostInfo(head))
assert.Equal(t, other, hm.QueryVpnAddr(a))
assert.Equal(t, other, hm.QueryVpnAddr(b))
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{1}, chainIds(t, hm, b))
// head is fully removed from the index map.
assert.Nil(t, hm.QueryIndex(head.localIndexId))
}
// TestHostMap_DeleteHostInfo_DivergentVpnAddrs covers chained hostinfos for the same peer whose
// vpnAddrs sets differ (a re-handshake cert added a second address). Deleting the superset node
// must not promote a sibling to an address it does not own.
func TestHostMap_DeleteHostInfo_DivergentVpnAddrs(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
b := netip.MustParseAddr("0.0.0.2")
// sub owns only a; super (a newer handshake) owns a and b.
sub := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
super := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 2}
hm.unlockedAddHostInfo(sub, f)
hm.unlockedAddHostInfo(super, f)
assert.Equal(t, []uint32{2, 1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{2}, chainIds(t, hm, b))
// Delete super: a promotes to sub (which owns it); b has no remaining owner and must be
// removed, not dangled at sub (which does not own b).
assert.False(t, hm.DeleteHostInfo(super))
assert.Equal(t, []uint32{1}, chainIds(t, hm, a))
assert.Empty(t, chainIds(t, hm, b))
assert.Equal(t, sub, hm.QueryVpnAddr(a))
assert.Nil(t, hm.QueryVpnAddr(b))
assert.Nil(t, hm.QueryIndex(super.localIndexId))
// Deleting sub cleans up fully.
assert.True(t, hm.DeleteHostInfo(sub))
assert.Nil(t, hm.QueryVpnAddr(a))
assertHostMapInvariants(t, hm)
}
// TestHostMap_AddDivergentOverlap covers a new hostinfo claiming addresses currently owned by two
// DIFFERENT hostinfos. The old single shared next/prev chain overwrote a pointer and orphaned one
// of them (in Indexes but unreachable via its address); independent per-address lists cannot.
func TestHostMap_AddDivergentOverlap(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
b := netip.MustParseAddr("0.0.0.2")
hiA := &HostInfo{vpnAddrs: []netip.Addr{a}, localIndexId: 1}
hiP := &HostInfo{vpnAddrs: []netip.Addr{b}, localIndexId: 2}
hm.unlockedAddHostInfo(hiA, f)
hm.unlockedAddHostInfo(hiP, f)
hiB := &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: 3}
hm.unlockedAddHostInfo(hiB, f)
assert.Equal(t, []uint32{3, 1}, chainIds(t, hm, a))
assert.Equal(t, []uint32{3, 2}, chainIds(t, hm, b))
// hiA is still reachable via its address (not orphaned) and still indexed.
assert.Contains(t, chainIds(t, hm, a), hiA.localIndexId)
assert.NotNil(t, hm.QueryIndex(hiA.localIndexId))
}
// TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs verifies the MaxHostInfosPerVpnIp overflow prune
// (unlockedInnerAddHostInfo calls unlockedDeleteHostInfo on the oldest node once the chain is too long)
// still behaves when hostinfos carry more than one vpnAddr. The pruned node is always the tail, so it is
// primary for none of the addresses, and both address chains must stay consistent afterwards.
func TestHostMap_MaxHostInfosPerVpnIp_MultipleVpnAddrs(t *testing.T) {
l := test.NewLogger()
hm := newHostMap(l)
f := &Interface{}
a := netip.MustParseAddr("0.0.0.1")
b := netip.MustParseAddr("0.0.0.2")
// Add one more than the cap, newest last so it becomes head. Every hostinfo owns both a and b.
hostinfos := make([]*HostInfo, 0, MaxHostInfosPerVpnIp+1)
for i := 0; i <= MaxHostInfosPerVpnIp; i++ {
hostinfos = append(hostinfos, &HostInfo{vpnAddrs: []netip.Addr{a, b}, localIndexId: uint32(i + 1)})
}
// Add oldest first (highest index in our slice) so the very first one added is the overflow victim.
for i := len(hostinfos) - 1; i >= 0; i-- {
hm.unlockedAddHostInfo(hostinfos[i], f)
}
oldest := hostinfos[len(hostinfos)-1]
// The oldest hostinfo was pruned from both lists and the index map.
assert.Nil(t, hm.QueryIndex(oldest.localIndexId))
// Both addresses hold exactly MaxHostInfosPerVpnIp survivors in the same order; oldest is absent.
require.Len(t, chainIds(t, hm, a), MaxHostInfosPerVpnIp)
assert.Equal(t, chainIds(t, hm, a), chainIds(t, hm, b), "both addresses must list the same survivors in the same order")
assert.NotContains(t, chainIds(t, hm, a), oldest.localIndexId)
assert.Equal(t, hm.QueryVpnAddr(a), hm.QueryVpnAddr(b))
} }
func TestHostMap_reload(t *testing.T) { func TestHostMap_reload(t *testing.T) {
+24 -26
View File
@@ -25,11 +25,11 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
// //
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO // pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
// superpacket. In both cases the L3+L4 headers at the start describe // superpacket. In both cases the L3+L4 headers at the start describe
// the same 5-tuple every segment will share, so a single newPacket / // the same 5-tuple every segment will share, so a single parse +
// firewall check covers the whole superpacket. // firewall check covers the whole superpacket.
packet := pkt.Bytes packet := pkt.Bytes
err := newPacket(packet, false, fwPacket) var parsed batch.RxParsed
if err != nil { if err := batch.ParsePacket(packet, false, &parsed); err != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("Error while validating outbound packet", f.l.Debug("Error while validating outbound packet",
"packet", packet, "packet", packet,
@@ -39,6 +39,8 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
return return
} }
parsed.Key.Hydrate(fwPacket)
// Ignore local broadcast packets // Ignore local broadcast packets
if f.dropLocalBroadcast { if f.dropLocalBroadcast {
if f.myBroadcastAddrsTable.Contains(fwPacket.RemoteAddr) { if f.myBroadcastAddrsTable.Contains(fwPacket.RemoteAddr) {
@@ -57,7 +59,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
// kernel as one giant blob; segment first so the loopback // kernel as one giant blob; segment first so the loopback
// path sees one IP datagram per Write. // path sees one IP datagram per Write.
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error { err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
_, werr := f.queues[q].Write(seg) _, werr := f.readers[q].Write(seg)
return werr return werr
}) })
if err != nil { if err != nil {
@@ -105,7 +107,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
return return
} }
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache) dropReason := f.firewall.Drop(parsed.Key, fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
if dropReason == nil { if dropReason == nil {
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q) f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
} else { } else {
@@ -135,7 +137,7 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
if encErr != nil { if encErr != nil {
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet", hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
"error", encErr, "error", encErr,
"udpAddr", hostinfo.GetRemote(), "udpAddr", hostinfo.remote,
"counter", c, "counter", c,
) )
// Skip this segment; the rest of the superpacket can still // Skip this segment; the rest of the superpacket can still
@@ -159,7 +161,6 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
return return
} }
remote := hostinfo.GetRemote()
ecnEnabled := f.ecnEnabled.Load() ecnEnabled := f.ecnEnabled.Load()
if hostinfo.lastRebindCount != f.rebindCount { if hostinfo.lastRebindCount != f.rebindCount {
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is //NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
@@ -173,7 +174,7 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
} }
} }
if !remote.IsValid() { //the relay path if !hostinfo.remote.IsValid() { //the relay path
//first, find our relay hostinfo: //first, find our relay hostinfo:
var relayHostInfo *HostInfo var relayHostInfo *HostInfo
var relay *Relay var relay *Relay
@@ -215,7 +216,7 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
if ecnEnabled { if ecnEnabled {
ecn = innerECN(seg) ecn = innerECN(seg)
} }
sendBatch.Commit(toSend, relayHostInfo.GetRemote(), ecn) sendBatch.Commit(toSend, relayHostInfo.remote, ecn)
return nil return nil
}) })
if err != nil { if err != nil {
@@ -237,7 +238,7 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
if ecnEnabled { if ecnEnabled {
ecn = innerECN(seg) ecn = innerECN(seg)
} }
sendBatch.Commit(out, remote, ecn) sendBatch.Commit(out, hostinfo.remote, ecn)
return nil return nil
}) })
if err != nil { if err != nil {
@@ -264,7 +265,7 @@ func innerECN(pkt []byte) byte {
} }
func (f *Interface) rejectInside(packet []byte, out []byte, q int) { func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
if !f.firewall.OutboundSendReject { if !f.firewall.InSendReject {
return return
} }
@@ -273,14 +274,14 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
return return
} }
_, err := f.queues[q].Write(out) _, err := f.readers[q].Write(out)
if err != nil { if err != nil {
f.l.Error("Failed to write to tun", "error", err) f.l.Error("Failed to write to tun", "error", err)
} }
} }
func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) { func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *HostInfo, nb, out []byte, q int) {
if !f.firewall.InboundSendReject { if !f.firewall.OutSendReject {
return return
} }
@@ -393,15 +394,16 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
} }
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) { func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
fp := &firewall.Packet{} var parsed batch.RxParsed
err := newPacket(p, false, fp) if err := batch.ParsePacket(p, false, &parsed); err != nil {
if err != nil {
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err) f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
return return
} }
fp := &firewall.Packet{}
parsed.Key.Hydrate(fp)
// check if packet is in outbound fw rules // check if packet is in outbound fw rules
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil) dropReason := f.firewall.Drop(parsed.Key, fp, false, hostinfo, f.pki.GetCAPool(), nil)
if dropReason != nil { if dropReason != nil {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
f.l.Debug("dropping cached packet", f.l.Debug("dropping cached packet",
@@ -522,12 +524,7 @@ func (f *Interface) SendVia(via *HostInfo,
nocopy bool, nocopy bool,
) { ) {
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy) toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
if err != nil { err = f.writers[0].WriteTo(toSend, via.remote)
// already logged by prepareSendVia
return
}
err = f.writers[0].WriteTo(toSend, via.GetRemote())
if err != nil { if err != nil {
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err) via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
} }
@@ -537,7 +534,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
if ci.eKey == nil { if ci.eKey == nil {
return return
} }
useRelay := !remote.IsValid() && !hostinfo.GetRemote().IsValid() useRelay := !remote.IsValid() && !hostinfo.remote.IsValid()
fullOut := out fullOut := out
if useRelay { if useRelay {
@@ -584,6 +581,7 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
"error", err, "error", err,
"udpAddr", remote, "udpAddr", remote,
"counter", c, "counter", c,
"attemptedCounter", c,
) )
return return
} }
@@ -596,8 +594,8 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
"udpAddr", remote, "udpAddr", remote,
) )
} }
} else if hr := hostinfo.GetRemote(); hr.IsValid() { } else if hostinfo.remote.IsValid() {
err = f.writers[q].WriteTo(out, hr) err = f.writers[q].WriteTo(out, hostinfo.remote)
if err != nil { if err != nil {
hostinfo.logger(f.l).Error("Failed to write outgoing packet", hostinfo.logger(f.l).Error("Failed to write outgoing packet",
"error", err, "error", err,
+58 -120
View File
@@ -7,7 +7,6 @@ import (
"log/slog" "log/slog"
"net/netip" "net/netip"
"runtime" "runtime"
"slices"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
@@ -16,7 +15,6 @@ import (
"github.com/rcrowley/go-metrics" "github.com/rcrowley/go-metrics"
"github.com/slackhq/nebula/util" "github.com/slackhq/nebula/util"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config" "github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/header" "github.com/slackhq/nebula/header"
@@ -56,13 +54,8 @@ type InterfaceConfig struct {
// CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine // CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
// should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] — // should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] —
// shorter lists than `routines` cycle. Empty list keeps the default // shorter lists than `routines` cycle. Empty list keeps the default
// pin-to-(i % NumCPU) behavior. Only consulted when PinThreads is true. // pin-to-(i % NumCPU) behavior.
CpuAffinity []int CpuAffinity []int
// PinThreads controls whether each TUN reader OS thread is pinned to a
// single CPU (via tun.pin_threads, default true). Pinning keeps each
// goroutine's sendmmsg on one XPS-selected NIC TX ring so per-flow
// packets stay ordered on the wire.
PinThreads bool
l *slog.Logger l *slog.Logger
} }
@@ -90,13 +83,8 @@ type Interface struct {
closed atomic.Bool closed atomic.Bool
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine // cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
// should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)]. // should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)].
// Empty falls back to the default pin-to-(allowed CPU) behavior. // Empty falls back to the default pin-to-(i % NumCPU) behavior.
// Only consulted when pinThreads is true.
cpuAffinity []int cpuAffinity []int
// pinThreads controls whether listenIn pins each TUN reader OS thread to
// a CPU at all (tun.pin_threads, default true). When false, threads are
// left free to migrate as on stock nebula.
pinThreads bool
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true, // ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
// inside.go copies the inner ECN onto the outer carrier on encap and // inside.go copies the inner ECN onto the outer carrier on encap and
// decryptToTun folds outer CE into the inner header on decap. Toggle // decryptToTun folds outer CE into the inner header on decap. Toggle
@@ -119,8 +107,8 @@ type Interface struct {
ctx context.Context ctx context.Context
writers []udp.Conn writers []udp.Conn
queues []tio.Queue readers []tio.Queue
// batchers is one per tun queue, wrapping queues[i]. // batchers is one per tun queue, wrapping readers[i].
// decryptToTun sends plaintext into the batch.RxBatcher; // decryptToTun sends plaintext into the batch.RxBatcher;
// listenOut calls its Flush at the end of each UDP recvmmsg batch. // listenOut calls its Flush at the end of each UDP recvmmsg batch.
batchers []batch.RxBatcher batchers []batch.RxBatcher
@@ -222,6 +210,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
routines: c.routines, routines: c.routines,
version: c.version, version: c.version,
writers: make([]udp.Conn, c.routines), writers: make([]udp.Conn, c.routines),
readers: make([]tio.Queue, c.routines),
batchers: make([]batch.RxBatcher, c.routines), batchers: make([]batch.RxBatcher, c.routines),
myVpnNetworks: cs.myVpnNetworks, myVpnNetworks: cs.myVpnNetworks,
myVpnNetworksTable: cs.myVpnNetworksTable, myVpnNetworksTable: cs.myVpnNetworksTable,
@@ -232,7 +221,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
connectionManager: c.connectionManager, connectionManager: c.connectionManager,
conntrackCacheTimeout: c.ConntrackCacheTimeout, conntrackCacheTimeout: c.ConntrackCacheTimeout,
cpuAffinity: c.CpuAffinity, cpuAffinity: c.CpuAffinity,
pinThreads: c.PinThreads,
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)), metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
messageMetrics: c.MessageMetrics, messageMetrics: c.MessageMetrics,
@@ -250,9 +238,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
ifce.connectionManager.intf = ifce ifce.connectionManager.intf = ifce
// Held until Close so waiting on the interface blocks until the resources are actually released
ifce.wg.Add(1)
return ifce, nil return ifce, nil
} }
@@ -275,52 +260,48 @@ func (f *Interface) activate() error {
"boringcrypto", boringEnabled(), "boringcrypto", boringEnabled(),
) )
if f.routines > 1 && !f.outside.SupportsMultipleReaders() { if f.routines > 1 {
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
f.routines = 1 f.routines = 1
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine") f.l.Warn("routines is not supported on this platform, falling back to a single routine")
} }
// Prepare the tun queues. A device that can't open that many hands back
// fewer (a single queue on platforms without multiqueue support) and we
// size the reader routines to what we actually got.
queues, err := f.inside.Queues(f.routines)
if err != nil {
return err
} }
if len(queues) < f.routines {
f.l.Warn("tun multiqueue is not supported on this platform, falling back to fewer routines",
"requested", f.routines, "opened", len(queues))
f.routines = len(queues)
}
f.queues = queues
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines)) metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
for i := range f.queues { // Prepare n tun queues
caps := tio.QueueCapabilities(f.queues[i]) for i := 0; i < f.routines; i++ {
if i > 0 {
if err = f.inside.NewMultiQueueReader(); err != nil {
return err
}
}
}
f.readers = f.inside.Readers()
for i := range f.readers {
caps := tio.QueueCapabilities(f.readers[i])
if caps.TSO || caps.USO { if caps.TSO || caps.USO {
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO // Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
// is on, everything else (and either lane disabled) falls // is on, everything else (and either lane disabled) falls
// through to passthrough so non-IP / non-TCP-UDP traffic still // through to passthrough so non-IP / non-TCP-UDP traffic still
// reaches the TUN. // reaches the TUN.
arena := batch.NewArena(batch.DefaultMultiArenaCap) f.batchers[i] = batch.NewMultiCoalescer(f.readers[i], caps.TSO, caps.USO)
f.batchers[i] = batch.NewMultiCoalescer(f.queues[i], f.l, arena, caps.TSO, caps.USO)
} else { } else {
arena := batch.NewArena(batch.DefaultPassthroughArenaCap) f.batchers[i] = batch.NewPassthrough(f.readers[i])
f.batchers[i] = batch.NewPassthrough(f.queues[i], arena.Reserve, arena.Reset)
} }
} }
// On error the caller owns the cleanup, Control.Start cancels the service context f.wg.Add(1) // for us to wait on Close() to return
// before releasing our resources so a waiter never observes a live context
if err = f.inside.Activate(); err != nil { if err = f.inside.Activate(); err != nil {
f.wg.Done()
f.inside.Close()
return err return err
} }
return nil return nil
} }
func (f *Interface) run() { func (f *Interface) run() (func() error, error) {
// Launch n queues to read packets from udp // Launch n queues to read packets from udp
for i := 0; i < f.routines; i++ { for i := 0; i < f.routines; i++ {
f.wg.Go(func() { f.wg.Go(func() {
@@ -331,18 +312,17 @@ func (f *Interface) run() {
// Launch n queues to read packets from tun dev // Launch n queues to read packets from tun dev
for i := 0; i < f.routines; i++ { for i := 0; i < f.routines; i++ {
f.wg.Go(func() { f.wg.Go(func() {
f.listenIn(f.queues[i], i) f.listenIn(f.readers[i], i)
}) })
} }
} return func() error {
func (f *Interface) wait() error {
f.wg.Wait() f.wg.Wait()
if e := f.fatalErr.Load(); e != nil { if e := f.fatalErr.Load(); e != nil {
return *e return *e
} }
return nil return nil
}, nil
} }
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one // onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
@@ -368,11 +348,12 @@ func (f *Interface) listenOut(i int) {
lhh := f.lightHouse.NewRequestHandler() lhh := f.lightHouse.NewRequestHandler()
h := &header.H{} h := &header.H{}
fwPacket := &firewall.Packet{} fwPacket := &firewall.Packet{}
parsedRx := &batch.RxParsed{}
nb := make([]byte, 12, 12) nb := make([]byte, 12, 12)
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) { listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
plaintext := f.batchers[i].Reserve(len(payload)) plaintext := f.batchers[i].Reserve(len(payload))
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(), meta) f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, parsedRx, lhh, nb, i, ctCache.Get(), meta)
} }
flusher := func() { flusher := func() {
@@ -383,10 +364,7 @@ func (f *Interface) listenOut(i int) {
err := li.ListenOut(listener, flusher) err := li.ListenOut(listener, flusher)
// An error after teardown began is shutdown noise, the closed flag covers resources if err != nil && !f.closed.Load() {
// Close releases itself and the cancelled ctx covers ones torn down by their owners
// reacting to it, like the user device pipes
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
f.l.Error("Error while reading inbound packet, closing", "error", err) f.l.Error("Error while reading inbound packet, closing", "error", err)
f.onFatal(err) f.onFatal(err)
} }
@@ -394,41 +372,37 @@ func (f *Interface) listenOut(i int) {
f.l.Debug("underlay reader is done", "reader", i) f.l.Debug("underlay reader is done", "reader", i)
} }
func (f *Interface) listenIn(queue tio.Queue, i int) { func (f *Interface) listenIn(reader tio.Queue, i int) {
// Pinning this thread (and goroutine) to a single CPU keeps every sendmmsg from this goroutine going through the // Pin this goroutine to one CPU. LockOSThread alone keeps the goroutine
// same TX ring on the nic, so the wire sees per-flow order. Skip entirely when tun.pin_threads is false. // on a single OS thread but the kernel can still migrate that thread
if f.pinThreads { // across CPUs — XPS reads smp_processor_id() at sendmmsg time and picks
var cpu int // the TX ring from the current CPU's xps_cpus map, so an unpinned
// thread bouncing between CPUs spreads one nebula flow's packets across
// multiple TX rings, which the rings then drain at independent rates
// and the wire delivers reordered.
//
// Pinning keeps every sendmmsg from this goroutine going through the
// same TX ring, so the wire sees per-flow order. Cost: less scheduler
// flexibility — if i % NumCPU collides between two TUN reader
// goroutines they share a CPU.
cpu := i % runtime.NumCPU()
if n := len(f.cpuAffinity); n > 0 { if n := len(f.cpuAffinity); n > 0 {
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
// validated the entries against the allowed CPU set.
cpu = f.cpuAffinity[i%n] cpu = f.cpuAffinity[i%n]
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
// Default: spread queues across the CPUs we're actually allowed to
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
// i % NumCPU would pick unrunnable IDs and every pin would fail.
cpu = allowed[i%len(allowed)]
} else {
cpu = i % runtime.NumCPU()
} }
if err := util.PinThreadToCPU(cpu); err != nil { if err := util.PinThreadToCPU(cpu); err != nil {
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err) f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
} }
}
rejectBuf := make([]byte, mtu) rejectBuf := make([]byte, mtu)
arenaSize := batch.SendBatchCap * (udp.MTU + 32) sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, udp.MTU+32)
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, arenaSize)
fwPacket := &firewall.Packet{} fwPacket := &firewall.Packet{}
nb := make([]byte, 12, 12) nb := make([]byte, 12, 12)
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout) conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
for { for {
pkts, err := queue.Read() pkts, err := reader.Read()
if err != nil { if err != nil {
// Same shutdown noise handling as listenOut if !f.closed.Load() {
if !f.closed.Load() && f.ctx.Err() == nil {
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i) f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
f.onFatal(err) f.onFatal(err)
} }
@@ -437,14 +411,6 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
for _, pkt := range pkts { for _, pkt := range pkts {
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get()) f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
// Flush incrementally once a full sendmmsg batch has
// accumulated so the first packets of a deep read drain
// hit the wire while the rest are still being encrypted.
if sb.Len() >= batch.SendBatchCap {
if err := sb.Flush(); err != nil {
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
}
}
} }
if err := sb.Flush(); err != nil { if err := sb.Flush(); err != nil {
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i) f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
@@ -478,22 +444,13 @@ func (f *Interface) reloadDisconnectInvalid(c *config.C) {
} }
func (f *Interface) reloadFirewall(c *config.C) { func (f *Interface) reloadFirewall(c *config.C) {
cs := f.pki.getCertState() //TODO: need to trigger/detect if the certificate changed too
curCert := cs.getCertificate(cert.Version2) if c.HasChanged("firewall") == false {
if curCert == nil {
curCert = cs.getCertificate(cert.Version1)
}
// The firewall builds its routableNetworks set from the certificate's UnsafeNetworks at construction.
// Check to see if that set has changed, and if so, rebuild the firewall.
certUnsafeChanged := curCert != nil && !slices.Equal(curCert.UnsafeNetworks(), f.firewall.unsafeNetworks)
if !c.HasChanged("firewall") && !certUnsafeChanged {
f.l.Debug("No firewall config change detected") f.l.Debug("No firewall config change detected")
return return
} }
fw, err := NewFirewallFromConfig(f.l, cs, c) fw, err := NewFirewallFromConfig(f.l, f.pki.getCertState(), c)
if err != nil { if err != nil {
f.l.Error("Error while creating firewall during reload", "error", err) f.l.Error("Error while creating firewall during reload", "error", err)
return return
@@ -600,12 +557,9 @@ func (f *Interface) reloadEcn(c *config.C) {
initial := c.InitialLoad() initial := c.InitialLoad()
if initial || c.HasChanged("tunnels.ecn") { if initial || c.HasChanged("tunnels.ecn") {
v := c.GetBool("tunnels.ecn", true) v := c.GetBool("tunnels.ecn", true)
changed := f.ecnEnabled.Swap(v) != v f.ecnEnabled.Store(v)
if !initial { if !initial {
f.l.Info("tunnels.ecn changed", "enabled", v) f.l.Info("tunnels.ecn changed", "enabled", v)
if changed {
f.l.Warn("tunnels.ecn datapath toggled, but route-level ECN negotiation (RTAX_FEATURE_ECN) retains its previous state until nebula is restarted", "enabled", v)
}
} }
} }
} }
@@ -620,7 +574,11 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil) certInitiatingVersion := metrics.GetOrRegisterGauge("certificate.initiating_version", nil)
certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil) certMaxVersion := metrics.GetOrRegisterGauge("certificate.max_version", nil)
emit := func() { for {
select {
case <-ctx.Done():
return
case <-ticker.C:
f.firewall.EmitStats() f.firewall.EmitStats()
f.handshakeManager.EmitStats() f.handshakeManager.EmitStats()
udpStats() udpStats()
@@ -637,18 +595,6 @@ func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
certMaxVersion.Update(int64(certState.v1Cert.Version())) certMaxVersion.Update(int64(certState.v1Cert.Version()))
} }
} }
// Prime gauges so a Prometheus scrape that lands before the first tick
// sees real values instead of the zero defaults (issue #907).
emit()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
emit()
}
} }
} }
@@ -660,15 +606,9 @@ func (f *Interface) GetCertState() *CertState {
return f.pki.getCertState() return f.pki.getCertState()
} }
// Close releases the interface's resources: the udp sockets and the tun device.
// It is idempotent and safe to call at any point in the lifecycle, including on an interface that never activated,
// calls after the first return nil without doing anything.
func (f *Interface) Close() error { func (f *Interface) Close() error {
if !f.closed.CompareAndSwap(false, true) {
return nil
}
var errs []error var errs []error
f.closed.Store(true)
// Release the udp readers // Release the udp readers
for i, u := range f.writers { for i, u := range f.writers {
@@ -684,8 +624,6 @@ func (f *Interface) Close() error {
if closeErr != nil { if closeErr != nil {
errs = append(errs, closeErr) errs = append(errs, closeErr)
} }
// Release the construction token so waiters know the resources are gone
f.wg.Done() f.wg.Done()
return errors.Join(errs...) return errors.Join(errs...)
} }
-73
View File
@@ -1,73 +0,0 @@
//go:build linux || darwin
package nebula
import (
"context"
"net/netip"
"testing"
"time"
"github.com/rcrowley/go-metrics"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/overlay/overlaytest"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/udp"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// Test_emitStats_primesGauges covers issue #907: a Prometheus scrape that
// landed before the first ticker fire used to read 0 for the cert gauges.
// emitStats now primes the gauges before entering the ticker loop. We assert
// the gauge is zero before the first call and non-zero after.
func Test_emitStats_primesGauges(t *testing.T) {
defer metrics.DefaultRegistry.UnregisterAll()
l := test.NewLogger()
hostMap := newHostMap(l)
preferredRanges := []netip.Prefix{netip.MustParsePrefix("10.0.0.0/8")}
hostMap.preferredRanges.Store(&preferredRanges)
notAfter := time.Now().Add(time.Hour)
cs := &CertState{
initiatingVersion: cert.Version1,
privateKey: []byte{},
v1Cert: &dummyCert{version: cert.Version1, notAfter: notAfter},
v1Credential: nil,
}
lh := newTestLighthouse()
ifce := &Interface{
hostMap: hostMap,
inside: &overlaytest.NoopTun{},
outside: &udp.NoopConn{},
firewall: &Firewall{Conntrack: &FirewallConntrack{Conns: map[firewall.Packet]*conn{}}},
lightHouse: lh,
pki: &PKI{},
handshakeManager: NewHandshakeManager(l, hostMap, lh, &udp.NoopConn{}, defaultHandshakeConfig),
l: l,
// On linux, udp.NewUDPStatsEmitter indexes writers[0] and asserts to
// *udp.StdConn. A zero value works: getMemInfo sees a nil rawConn,
// returns an error, and the emitter falls through to a no-op.
writers: []udp.Conn{&udp.StdConn{}},
}
ifce.pki.cs.Store(cs)
ttlGauge := metrics.GetOrRegisterGauge("certificate.ttl_seconds", nil)
require.Zero(t, ttlGauge.Value(), "gauge should be zero before emitStats runs")
// Pre-cancel the context so emitStats returns after priming the gauges
// without ever reading from ticker.C. The one hour interval is just a
// belt-and-suspenders, the test does not expect the ticker to fire.
ctx, cancel := context.WithCancel(context.Background())
cancel()
ifce.emitStats(ctx, time.Hour)
ttl := ttlGauge.Value()
assert.Positive(t, ttl, "ttl gauge should be primed by emitStats before its first tick")
assert.LessOrEqual(t, ttl, int64(3600))
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.initiating_version", nil).Value())
assert.Equal(t, int64(cert.Version1), metrics.GetOrRegisterGauge("certificate.max_version", nil).Value())
}
-120
View File
@@ -1,120 +0,0 @@
package nebula
import (
"net/netip"
"testing"
"github.com/slackhq/nebula/cert"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// TestReloadFirewall_CertUnsafeNetworksChanged verifies that reloadFirewall
// rebuilds the firewall when only the certificate's UnsafeNetworks have changed,
// even if the firewall section of the YAML has not.
func TestReloadFirewall_CertUnsafeNetworksChanged(t *testing.T) {
l := test.NewLogger()
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
initialUnsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
// dummyCert avoids dragging the real signing pipeline into a unit test.
c1 := &dummyCert{
version: cert.Version2,
networks: []netip.Prefix{vpnNet},
unsafeNetworks: initialUnsafe,
}
pki := &PKI{}
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
rawYAML := `firewall:
outbound:
- port: any
proto: any
host: any
inbound:
- port: any
proto: any
host: any
`
cfg := config.NewC(l)
require.NoError(t, cfg.LoadString(rawYAML))
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
require.NoError(t, err)
require.Equal(t, initialUnsafe, fw.unsafeNetworks)
f := &Interface{
pki: pki,
firewall: fw,
l: l,
}
// Swap the cert with a different UnsafeNetworks set.
newUnsafe := []netip.Prefix{
netip.MustParsePrefix("198.51.100.0/24"),
netip.MustParsePrefix("203.0.113.0/24"),
}
c2 := &dummyCert{
version: cert.Version2,
networks: []netip.Prefix{vpnNet},
unsafeNetworks: newUnsafe,
}
pki.cs.Store(&CertState{v2Cert: c2, initiatingVersion: cert.Version2})
// Reload with the same YAML so HasChanged("firewall") reports false.
require.NoError(t, cfg.ReloadConfigString(rawYAML))
require.False(t, cfg.HasChanged("firewall"))
f.reloadFirewall(cfg)
assert.NotSame(t, fw, f.firewall, "firewall pointer should have been replaced")
assert.Equal(t, newUnsafe, f.firewall.unsafeNetworks)
assert.True(t, f.firewall.routableNetworks.Contains(netip.MustParseAddr("203.0.113.5")))
}
// TestReloadFirewall_NoChange verifies that reloadFirewall is a no-op when
// neither the firewall config nor the cert's UnsafeNetworks have changed.
func TestReloadFirewall_NoChange(t *testing.T) {
l := test.NewLogger()
vpnNet := netip.MustParsePrefix("10.0.0.1/24")
unsafe := []netip.Prefix{netip.MustParsePrefix("198.51.100.0/24")}
c1 := &dummyCert{
version: cert.Version2,
networks: []netip.Prefix{vpnNet},
unsafeNetworks: unsafe,
}
pki := &PKI{}
pki.cs.Store(&CertState{v2Cert: c1, initiatingVersion: cert.Version2})
rawYAML := `firewall:
outbound:
- port: any
proto: any
host: any
inbound:
- port: any
proto: any
host: any
`
cfg := config.NewC(l)
require.NoError(t, cfg.LoadString(rawYAML))
fw, err := NewFirewallFromConfig(l, pki.getCertState(), cfg)
require.NoError(t, err)
f := &Interface{
pki: pki,
firewall: fw,
l: l,
}
require.NoError(t, cfg.ReloadConfigString(rawYAML))
f.reloadFirewall(cfg)
assert.Same(t, fw, f.firewall, "firewall should not have been replaced")
}
+8 -279
View File
@@ -4,55 +4,27 @@ import (
"encoding/binary" "encoding/binary"
"golang.org/x/net/ipv4" "golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
) )
const ( const (
// MaxIPv4RejectPacketSize is the largest IPv4 reject packet: // Need 96 bytes for the largest reject packet:
// - 20 byte ipv4 header // - 20 byte ipv4 header
// - 8 byte icmpv4 header // - 8 byte icmpv4 header
// - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header) // - 68 byte body (60 byte max orig ipv4 header + 8 byte orig icmpv4 header)
maxIPv4RejectPacketSize = ipv4.HeaderLen + 8 + 60 + 8 MaxRejectPacketSize = 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 { func CreateRejectPacket(packet []byte, out []byte) []byte {
if len(packet) < 1 { if len(packet) < ipv4.HeaderLen || int(packet[0]>>4) != ipv4.Version {
return nil return nil
} }
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] { switch packet[9] {
case 6: // tcp case 6: // tcp
return ipv4CreateRejectTCPPacket(packet, out) return ipv4CreateRejectTCPPacket(packet, out)
default: default:
return ipv4CreateRejectICMPPacket(packet, out) return ipv4CreateRejectICMPPacket(packet, out)
} }
case ipv6.Version:
if len(packet) < ipv6.HeaderLen {
return nil
}
return ipv6CreateRejectPacket(packet, out)
default:
return nil
}
} }
func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte { func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
@@ -63,16 +35,11 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
return nil return nil
} }
// 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 // ICMP reply includes original header and first 8 bytes of the packet
packetLen := min(len(packet), ihl+8) packetLen := len(packet)
if packetLen > ihl+8 {
packetLen = ihl + 8
}
outLen := ipv4.HeaderLen + 8 + packetLen outLen := ipv4.HeaderLen + 8 + packetLen
if outLen > cap(out) { if outLen > cap(out) {
@@ -105,7 +72,7 @@ func ipv4CreateRejectICMPPacket(packet []byte, out []byte) []byte {
// ICMP Destination Unreachable // ICMP Destination Unreachable
icmpOut := out[ipv4.HeaderLen:] icmpOut := out[ipv4.HeaderLen:]
icmpOut[0] = 3 // type (Destination unreachable) icmpOut[0] = 3 // type (Destination unreachable)
icmpOut[1] = 13 // code (Communication administratively prohibited) icmpOut[1] = 3 // code (Port unreachable error)
icmpOut[2] = 0 // checksum icmpOut[2] = 0 // checksum
icmpOut[3] = 0 // . icmpOut[3] = 0 // .
icmpOut[4] = 0 // unused icmpOut[4] = 0 // unused
@@ -198,193 +165,7 @@ func ipv4CreateRejectTCPPacket(packet []byte, out []byte) []byte {
return out 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 { 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 // Return early if this is not a simple ICMP Echo Request
//TODO: make constants out of these //TODO: make constants out of these
if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) { if !(len(packet) >= 28 && len(packet) <= 9001 && packet[0] == 0x45 && packet[9] == 0x01 && packet[20] == 0x08) {
@@ -418,43 +199,6 @@ func createICMPv4EchoResponse(packet, out []byte) []byte {
return out 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 // calculates the TCP/IP checksum defined in rfc1071. The passed-in
// csum is any initial checksum data that's already been computed. // csum is any initial checksum data that's already been computed.
// //
@@ -492,18 +236,3 @@ func ipv4PseudoheaderChecksum(src, dst []byte, proto, length uint32) (csum uint3
csum += length >> 16 csum += length >> 16
return csum 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
}
+1 -445
View File
@@ -1,14 +1,11 @@
package iputil package iputil
import ( import (
"bytes"
"encoding/binary"
"net" "net"
"testing" "testing"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"golang.org/x/net/ipv4" "golang.org/x/net/ipv4"
"golang.org/x/net/ipv6"
) )
func Test_CreateRejectPacket(t *testing.T) { func Test_CreateRejectPacket(t *testing.T) {
@@ -46,7 +43,7 @@ func Test_CreateRejectPacket(t *testing.T) {
} }
b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...) b = append(b, []byte{0, 3, 0, 4, 0, 0, 0, 0}...)
expectedLen = maxIPv4RejectPacketSize expectedLen = MaxRejectPacketSize
out = make([]byte, MaxRejectPacketSize) out = make([]byte, MaxRejectPacketSize)
rejectPacket = CreateRejectPacket(b, out) rejectPacket = CreateRejectPacket(b, out)
assert.NotNil(t, rejectPacket) assert.NotNil(t, rejectPacket)
@@ -74,444 +71,3 @@ func Test_CreateRejectPacket(t *testing.T) {
assert.NotNil(t, rejectPacket) assert.NotNil(t, rejectPacket)
assert.Len(t, rejectPacket, expectedLen) 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)
}
}
// Test_CreateRejectPacket_RespectsCap ensures it is impossible for
// an oversized ICMPv6 reject to overwrite the neighbor segment's bytes.
func Test_CreateRejectPacket_RespectsCap(t *testing.T) {
src := net.ParseIP("fd00::1")
dst := net.ParseIP("fd00::2")
// Inner IPv6 UDP packet. An ICMPv6 reject copies the whole inner packet
// plus a 48-byte header (40 IPv6 + 8 ICMPv6), so it needs 48 more bytes
// than the inner packet length.
inner := makeIPv6Packet(src, dst, 17, make([]byte, 20))
// The ciphertext scratch reused as the reject buffer is the received
// datagram: 16-byte Nebula header + inner + 16-byte AEAD tag. That is only
// 32 bytes of slack, so a full ICMPv6 reject overruns it by 16 bytes.
const nebulaOverhead = 32
segLen := len(inner) + nebulaOverhead
// Shared backing row laid out as [segment][neighbor's 16-byte Nebula header].
const neighborHdr = 16
sentinel := bytes.Repeat([]byte{0xAB}, neighborHdr)
// Uncapped: the slice's capacity reaches into the neighbor, reproducing
// the overrun that silently drops the neighbor packet.
backing := make([]byte, segLen+neighborHdr)
copy(backing[segLen:], sentinel)
reject := CreateRejectPacket(inner, backing[:segLen])
assert.NotNil(t, reject, "uncapped buffer reaches into the neighbor, so the reject is built")
assert.NotEqual(t, sentinel, backing[segLen:segLen+neighborHdr],
"without the cap the oversized reject overruns into the neighbor segment")
// Capped (the fix): cap==len, so the builder cannot exceed the segment. The
// reject does not fit, so it is refused rather than corrupting the neighbor.
backing = make([]byte, segLen+neighborHdr)
copy(backing[segLen:], sentinel)
reject = CreateRejectPacket(inner, backing[:segLen:segLen])
assert.Nil(t, reject, "capped segment is 16 bytes too small for a full ICMPv6 reject, so it is refused")
assert.Equal(t, sentinel, backing[segLen:segLen+neighborHdr],
"capped segment must leave the neighbor untouched")
}
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)
}
+5 -30
View File
@@ -272,18 +272,16 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
//NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config //NOTE: many things will get much simpler when we combine static_host_map and lighthouse.hosts in config
if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") { if initial || c.HasChanged("static_host_map") || c.HasChanged("static_map.cadence") || c.HasChanged("static_map.network") || c.HasChanged("static_map.lookup_timeout") {
// Clean up. Entries still in the static_host_map will be re-built. // Clean up. Entries still in the static_host_map will be re-built.
ourselves := lh.myVpnNetworks[0].Addr() // Entries no longer present must have their (possible) background DNS goroutines stopped.
oldStaticList := lh.staticList.Load() if existingStaticList := lh.staticList.Load(); existingStaticList != nil {
if oldStaticList != nil {
lh.RLock() lh.RLock()
for staticVpnAddr := range *oldStaticList { for staticVpnAddr := range *existingStaticList {
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil { if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
am.ResetForOwner(ourselves) am.hr.Cancel()
} }
} }
lh.RUnlock() lh.RUnlock()
} }
// Build a new list based on current config. // Build a new list based on current config.
staticList := make(map[netip.Addr]struct{}) staticList := make(map[netip.Addr]struct{})
err := lh.loadStaticMap(c, staticList) err := lh.loadStaticMap(c, staticList)
@@ -291,21 +289,6 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
return err return err
} }
// For entries removed from static_host_map, stop the DNS goroutine and drop the cached addrs.
// All addrs must come from the lighthouses now that it's no longer a static host.
if oldStaticList != nil {
lh.RLock()
for staticVpnAddr := range *oldStaticList {
if _, stillStatic := staticList[staticVpnAddr]; stillStatic {
continue
}
if am, ok := lh.addrMap[staticVpnAddr]; ok && am != nil {
am.ClearHostnameResults()
}
}
lh.RUnlock()
}
lh.staticList.Store(&staticList) lh.staticList.Store(&staticList)
if !initial { if !initial {
if c.HasChanged("static_host_map") { if c.HasChanged("static_host_map") {
@@ -1418,9 +1401,6 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn
remoteAllowList := lhh.lh.GetRemoteAllowList() remoteAllowList := lhh.lh.GetRemoteAllowList()
for _, a := range n.Details.V4AddrPorts { for _, a := range n.Details.V4AddrPorts {
if a == nil {
continue
}
b := protoV4AddrPortToNetAddrPort(a) b := protoV4AddrPortToNetAddrPort(a)
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) { if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
lhh.lh.punchy.Schedule(b, detailsVpnAddr) lhh.lh.punchy.Schedule(b, detailsVpnAddr)
@@ -1428,9 +1408,6 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn
} }
for _, a := range n.Details.V6AddrPorts { for _, a := range n.Details.V6AddrPorts {
if a == nil {
continue
}
b := protoV6AddrPortToNetAddrPort(a) b := protoV6AddrPortToNetAddrPort(a)
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) { if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
lhh.lh.punchy.Schedule(b, detailsVpnAddr) lhh.lh.punchy.Schedule(b, detailsVpnAddr)
@@ -1460,7 +1437,7 @@ func protoV6AddrPortToNetAddrPort(ap *V6AddrPort) netip.AddrPort {
b := [16]byte{} b := [16]byte{}
binary.BigEndian.PutUint64(b[:8], ap.Hi) binary.BigEndian.PutUint64(b[:8], ap.Hi)
binary.BigEndian.PutUint64(b[8:], ap.Lo) binary.BigEndian.PutUint64(b[8:], ap.Lo)
return netip.AddrPortFrom(netip.AddrFrom16(b).Unmap(), uint16(ap.Port)) return netip.AddrPortFrom(netip.AddrFrom16(b), uint16(ap.Port))
} }
func netAddrToProtoAddr(addr netip.Addr) *Addr { func netAddrToProtoAddr(addr netip.Addr) *Addr {
@@ -1500,11 +1477,9 @@ func (d *NebulaMetaDetails) GetRelays() []netip.Addr {
if len(d.RelayVpnAddrs) > 0 { if len(d.RelayVpnAddrs) > 0 {
for _, r := range d.RelayVpnAddrs { for _, r := range d.RelayVpnAddrs {
if r != nil {
relays = append(relays, protoAddrToNetAddr(r)) relays = append(relays, protoAddrToNetAddr(r))
} }
} }
}
return relays return relays
} }
-126
View File
@@ -303,132 +303,6 @@ func TestLighthouse_reload(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
} }
// TestLighthouse_reloadStaticHostMap verifies that reloading static_host_map applies the new
// config rather than appending to it. See issue #718.
func TestLighthouse_reloadStaticHostMap(t *testing.T) {
l := test.NewLogger()
c := config.NewC(l)
c.Settings["lighthouse"] = map[string]any{"am_lighthouse": true}
c.Settings["listen"] = map[string]any{"port": 4242}
c.Settings["static_host_map"] = map[string]any{
"10.128.0.2": []any{"1.1.1.1:4242"},
}
myVpnNet := netip.MustParsePrefix("10.128.0.1/24")
nt := new(bart.Lite)
nt.Insert(myVpnNet)
cs := &CertState{
myVpnNetworks: []netip.Prefix{myVpnNet},
myVpnNetworksTable: nt,
}
lh, err := NewLightHouseFromConfig(t.Context(), l, c, cs, nil, nil)
require.NoError(t, err)
staticHost := netip.MustParseAddr("10.128.0.2")
otherHost := netip.MustParseAddr("10.128.0.3")
// Capture the RemoteList pointer up front; an in-flight handshake would hold the same one
// on hostinfo.remotes, so it must reflect every reload below.
pinned := lh.Query(staticHost)
require.NotNil(t, pinned)
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, pinned.CopyAddrs([]netip.Prefix{}))
// Replace the remote address. The new address should be the only entry.
nc := map[string]any{
"static_host_map": map[string]any{
"10.128.0.2": []any{"2.2.2.2:4242"},
},
}
rc, err := yaml.Marshal(nc)
require.NoError(t, err)
require.NoError(t, c.ReloadConfigString(string(rc)))
rl := lh.Query(staticHost)
require.NotNil(t, rl)
assert.Same(t, pinned, rl, "RemoteList pointer must stay stable so in-flight handshakes pick up the change")
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
// Reload back to the original IP. Mirrors the round-trip in issue #718 step 6-8 where
// the buggy reload produced [1.1.1.1, 2.2.2.2, 1.1.1.1] instead of [1.1.1.1].
nc = map[string]any{
"static_host_map": map[string]any{
"10.128.0.2": []any{"1.1.1.1:4242"},
},
}
rc, err = yaml.Marshal(nc)
require.NoError(t, err)
require.NoError(t, c.ReloadConfigString(string(rc)))
rl = lh.Query(staticHost)
require.NotNil(t, rl)
assert.Same(t, pinned, rl)
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
// Reload with the same config. An unchanged entry must not duplicate.
require.NoError(t, c.ReloadConfigString(string(rc)))
rl = lh.Query(staticHost)
require.NotNil(t, rl)
assert.Same(t, pinned, rl)
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("1.1.1.1:4242")}, rl.CopyAddrs([]netip.Prefix{}))
// Switch back to 2.2.2.2 so the rest of the test continues against a known address.
nc = map[string]any{
"static_host_map": map[string]any{
"10.128.0.2": []any{"2.2.2.2:4242"},
},
}
rc, err = yaml.Marshal(nc)
require.NoError(t, err)
require.NoError(t, c.ReloadConfigString(string(rc)))
// Add a second host alongside the first. Both should be present, neither duplicated.
nc = map[string]any{
"static_host_map": map[string]any{
"10.128.0.2": []any{"2.2.2.2:4242"},
"10.128.0.3": []any{"3.3.3.3:4242"},
},
}
rc, err = yaml.Marshal(nc)
require.NoError(t, err)
require.NoError(t, c.ReloadConfigString(string(rc)))
rl = lh.Query(staticHost)
require.NotNil(t, rl)
assert.Same(t, pinned, rl, "adding a sibling entry must not displace the existing RemoteList")
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("2.2.2.2:4242")}, rl.CopyAddrs([]netip.Prefix{}))
rl = lh.Query(otherHost)
require.NotNil(t, rl)
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
// Drop the first host entirely. The vpnAddr is no longer marked static, our owner
// contribution is cleared, but the addrMap entry stays in place so non-static cache
// data (from lighthouse queries) on the same RemoteList isn't lost. In-flight handshakes
// that already had the pointer see an empty address list rather than retrying stale ones.
nc = map[string]any{
"static_host_map": map[string]any{
"10.128.0.3": []any{"3.3.3.3:4242"},
},
}
rc, err = yaml.Marshal(nc)
require.NoError(t, err)
require.NoError(t, c.ReloadConfigString(string(rc)))
_, isStatic := lh.GetStaticHostList()[staticHost]
assert.False(t, isStatic)
rl = lh.Query(staticHost)
require.NotNil(t, rl)
assert.Same(t, pinned, rl)
assert.Empty(t, rl.CopyAddrs([]netip.Prefix{}))
rl = lh.Query(otherHost)
require.NotNil(t, rl)
assert.Equal(t, []netip.AddrPort{netip.MustParseAddrPort("3.3.3.3:4242")}, rl.CopyAddrs([]netip.Prefix{}))
}
func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply { func newLHHostRequest(fromAddr netip.AddrPort, myVpnIp, queryVpnIp netip.Addr, lhh *LightHouseHandler) testLhReply {
req := &NebulaMeta{ req := &NebulaMeta{
Type: NebulaMeta_HostQuery, Type: NebulaMeta_HostQuery,
+15 -99
View File
@@ -5,9 +5,11 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"net" "net"
"net/http"
_ "net/http/pprof"
"net/netip" "net/netip"
"runtime"
"runtime/debug" "runtime/debug"
"slices"
"strings" "strings"
"time" "time"
@@ -34,8 +36,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
buildVersion = moduleVersion() buildVersion = moduleVersion()
} }
// Debug builds (-tags debug) serve pprof on :6060; a no-op otherwise. //todo no merge
startPprofServer(ctx, l) go http.ListenAndServe(":6060", nil)
// Print the config if in test, the exit comes later // Print the config if in test, the exit comes later
if configTest { if configTest {
@@ -134,17 +136,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
udpConns := make([]udp.Conn, routines) udpConns := make([]udp.Conn, routines)
port := c.GetInt("listen.port", 0) port := c.GetInt("listen.port", 0)
// Callers get no handle to these until the Control is returned, release them on any error.
defer func() {
if reterr != nil {
for _, u := range udpConns {
if u != nil {
_ = u.Close()
}
}
}
}()
if !configTest { if !configTest {
rawListenHost := c.GetString("listen.host", "0.0.0.0") rawListenHost := c.GetString("listen.host", "0.0.0.0")
var listenHost netip.Addr var listenHost netip.Addr
@@ -209,17 +200,11 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig) handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
lightHouse.handshakeTrigger = handshakeManager.trigger lightHouse.handshakeTrigger = handshakeManager.trigger
ds, err := newDnsServerFromConfig(ctx, l, pki, hostMap, c) ds, err := newDnsServerFromConfig(ctx, l, pki.getCertState(), hostMap, c)
if err != nil { if err != nil {
l.Warn("Failed to start DNS responder", "error", err) l.Warn("Failed to start DNS responder", "error", err)
} }
pinThreads := c.GetBool("tun.pin_threads", true)
cpuAffinity := parseCpuAffinity(c, l, routines)
if pinThreads && len(cpuAffinity) == 0 && !configTest {
cpuAffinity = defaultCPUAffinityAvoidingIRQs(l, routines)
}
ifConfig := &InterfaceConfig{ ifConfig := &InterfaceConfig{
HostMap: hostMap, HostMap: hostMap,
Inside: tun, Inside: tun,
@@ -241,8 +226,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
relayManager: NewRelayManager(ctx, l, hostMap, c), relayManager: NewRelayManager(ctx, l, hostMap, c),
punchy: punchy, punchy: punchy,
ConntrackCacheTimeout: conntrackCacheTimeout, ConntrackCacheTimeout: conntrackCacheTimeout,
CpuAffinity: cpuAffinity, CpuAffinity: parseCpuAffinity(c, l, routines),
PinThreads: pinThreads,
l: l, l: l,
} }
@@ -297,16 +281,11 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of // parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of
// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil // integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil
// (listenIn falls back to spreading queues across the allowed CPU set). // (listenIn falls back to its default `i % NumCPU` pinning). Length
// Length mismatch with `routines` is a warning, not an error: shorter lists // mismatch with `routines` is a warning, not an error: shorter lists are
// are modulo-cycled across queues, longer lists' tail is ignored. Invalid // modulo-cycled across queues, longer lists' tail is ignored. Invalid
// entries (non-integer, or a CPU ID we're not allowed to run on) are also a // entries (non-integer, out of range) are also a warning and disable the
// warning and disable the override entirely so we don't silently pin to the // override entirely so we don't silently pin to the wrong CPU.
// wrong CPU. Entries are validated against the process's current affinity
// mask (util.AllowedCPUs) rather than 0..NumCPU-1: under a cgroup cpuset or
// taskset the runnable IDs are frequently not that contiguous range, and
// pinning to an unrunnable ID always fails. If the allowed set can't be
// determined we fall back to a plain non-negative check.
func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int { func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
raw := c.Get("tun.cpu_affinity") raw := c.Get("tun.cpu_affinity")
if raw == nil { if raw == nil {
@@ -317,14 +296,7 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw) l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw)
return nil return nil
} }
// allowed is the set of CPU IDs we're actually permitted to run on. A nil nCPU := runtime.NumCPU()
// slice (unsupported platform or lookup error) means "can't tell", so we
// only apply the weaker non-negative check in that case.
allowed, err := util.AllowedCPUs()
if err != nil {
l.Warn("could not determine allowed CPUs; validating tun.cpu_affinity against non-negative only", "error", err)
allowed = nil
}
cpus := make([]int, 0, len(rv)) cpus := make([]int, 0, len(rv))
for i, e := range rv { for i, e := range rv {
var cpu int var cpu int
@@ -340,14 +312,9 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
"index", i, "value", e) "index", i, "value", e)
return nil return nil
} }
if cpu < 0 { if cpu < 0 || cpu >= nCPU {
l.Warn("tun.cpu_affinity entry out of range; ignoring affinity", l.Warn("tun.cpu_affinity entry out of range; ignoring affinity",
"index", i, "cpu", cpu) "index", i, "cpu", cpu, "num_cpu", nCPU)
return nil
}
if len(allowed) > 0 && !slices.Contains(allowed, cpu) {
l.Warn("tun.cpu_affinity entry not in allowed CPU set; ignoring affinity",
"index", i, "cpu", cpu, "allowed", allowed)
return nil return nil
} }
cpus = append(cpus, cpu) cpus = append(cpus, cpu)
@@ -359,57 +326,6 @@ func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
return cpus return cpus
} }
// defaultCPUAffinityAvoidingIRQs picks the default pin set for the tun
// readers when tun.cpu_affinity is unset: allowed CPUs that do NOT service
// any physical NIC's interrupts. The stock allowed[i] spread pins the
// encrypt threads onto exactly the cores most drivers affine their first RX
// queue IRQs to, so whenever a flow's RSS queue fires on a core hosting a
// tun reader, NAPI and encrypt fight for the core and per-flow throughput
// drops (measured: REV 8.4 vs 10.2 Gbps on the same hardware, 2026-07-14).
//
// Returns nil — keeping the old allowed[i] fallback in listenIn — when IRQ
// info is unavailable or when there aren't enough IRQ-free CPUs to give
// every routine its own core: silently doubling readers up on fewer cores
// is worse than the occasional IRQ collision. NICs whose vectors blanket
// every CPU (e.g. mlx5 defaults to one queue per core) make avoidance
// impossible; narrowing the NIC's spread (ethtool -X <dev> equal N, or
// /proc/irq/*/smp_affinity) or setting tun.cpu_affinity explicitly makes it
// effective.
func defaultCPUAffinityAvoidingIRQs(l *slog.Logger, routines int) []int {
irq, err := util.NICIRQCPUs()
if err != nil || len(irq) == 0 {
return nil
}
allowed, err := util.AllowedCPUs()
if err != nil {
return nil
}
cpus := chooseIRQFreeCPUs(allowed, irq, routines)
if cpus == nil {
l.Info("not enough CPUs are free of NIC IRQs to give every tun reader its own; using the default spread",
"routines", routines, "allowed", len(allowed), "irqCPUs", len(irq))
return nil
}
l.Info("pinning tun readers to CPUs clear of NIC IRQs", "cpus", cpus)
return cpus
}
// chooseIRQFreeCPUs returns the first `routines` allowed CPUs not present in
// irq, or nil if fewer than `routines` qualify.
func chooseIRQFreeCPUs(allowed []int, irq map[int]bool, routines int) []int {
free := make([]int, 0, routines)
for _, cpu := range allowed {
if irq[cpu] {
continue
}
free = append(free, cpu)
if len(free) == routines {
return free
}
}
return nil
}
func moduleVersion() string { func moduleVersion() string {
info, ok := debug.ReadBuildInfo() info, ok := debug.ReadBuildInfo()
if !ok { if !ok {
-71
View File
@@ -1,71 +0,0 @@
package nebula
import (
"testing"
"github.com/slackhq/nebula/config"
"github.com/slackhq/nebula/test"
"github.com/slackhq/nebula/util"
"github.com/stretchr/testify/assert"
)
func TestChooseIRQFreeCPUs(t *testing.T) {
irq := map[int]bool{0: true, 1: true, 2: true, 3: true}
// Plenty of IRQ-free CPUs: take the first `routines` of them in order.
assert.Equal(t, []int{4, 5}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 2))
// Exactly enough.
assert.Equal(t, []int{4, 5, 6}, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4, 5, 6}, irq, 3))
// Not enough IRQ-free CPUs: nil, caller keeps the old default rather
// than doubling readers up on shared cores.
assert.Nil(t, chooseIRQFreeCPUs([]int{0, 1, 2, 3, 4}, irq, 2))
// No IRQ info at all behaves like a plain prefix of allowed.
assert.Equal(t, []int{0, 1}, chooseIRQFreeCPUs([]int{0, 1, 2}, map[int]bool{}, 2))
// Non-contiguous allowed set (cgroup cpuset) with holes.
assert.Equal(t, []int{9, 12}, chooseIRQFreeCPUs([]int{1, 3, 9, 12}, map[int]bool{1: true, 3: true}, 2))
}
func TestParseCpuAffinity(t *testing.T) {
l := test.NewLogger()
// newConfig returns a config.C with tun.cpu_affinity set to v. A nil v
// leaves the key unset.
newConfig := func(v any) *config.C {
c := config.NewC(l)
if v != nil {
c.Settings["tun"] = map[string]any{"cpu_affinity": v}
}
return c
}
// unset -> nil (listenIn falls back to spreading across the allowed set)
assert.Nil(t, parseCpuAffinity(newConfig(nil), l, 1))
// Pick a CPU we're actually allowed to run on so a valid list survives
// validation regardless of the host's affinity mask.
allowed, _ := util.AllowedCPUs()
validCPU := 0
if len(allowed) > 0 {
validCPU = allowed[0]
}
// valid list -> parsed through unchanged
assert.Equal(t, []int{validCPU, validCPU}, parseCpuAffinity(newConfig([]any{validCPU, validCPU}), l, 2))
// a negative entry is out of range on every platform -> disables the override
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, -1}), l, 2))
// a non-integer entry -> disables the override
assert.Nil(t, parseCpuAffinity(newConfig([]any{validCPU, "not-a-cpu"}), l, 2))
// a CPU id outside the allowed set -> disables the override. Only assertable
// where we can enumerate the allowed set (e.g. linux); 1<<20 is far beyond
// any representable CPU id so it can never be in the mask.
if len(allowed) > 0 {
assert.Nil(t, parseCpuAffinity(newConfig([]any{1 << 20}), l, 1))
}
}
+38 -238
View File
@@ -2,28 +2,20 @@ package nebula
import ( import (
"context" "context"
"encoding/binary"
"errors" "errors"
"log/slog" "log/slog"
"net/netip" "net/netip"
"time" "time"
"github.com/google/gopacket/layers"
"golang.org/x/net/ipv6"
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/header" "github.com/slackhq/nebula/header"
"github.com/slackhq/nebula/overlay/batch"
"github.com/slackhq/nebula/udp" "github.com/slackhq/nebula/udp"
"golang.org/x/net/ipv4"
)
const (
minFwPacketLen = 4
) )
var ErrOutOfWindow = errors.New("out of window packet") var ErrOutOfWindow = errors.New("out of window packet")
func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) { func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, parsedRx *batch.RxParsed, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
err := h.Parse(packet) err := h.Parse(packet)
if err != nil { if err != nil {
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors // Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
@@ -111,7 +103,8 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
// Relay packets are special // Relay packets are special
if isMessageRelay { if isMessageRelay {
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, lhf, nb, q, localCache, meta) f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, parsedRx, lhf, nb, q, localCache, meta)
return return
} }
@@ -135,7 +128,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
case header.Message: case header.Message:
switch h.Subtype { switch h.Subtype {
case header.MessageNone: case header.MessageNone:
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache, meta) f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, parsedRx, nb, q, localCache, meta)
default: default:
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h) hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
return return
@@ -150,8 +143,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
case header.TestReply: case header.TestReply:
// No-op, useful for the Roaming and connectionManager side-effects above // No-op, useful for the Roaming and connectionManager side-effects above
case header.TestRequest: case header.TestRequest:
//recycle the input packet ciphertext as our output buffer f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, out)
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, packet)
default: default:
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h) hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
return return
@@ -169,7 +161,7 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
} }
} }
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) { func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, out []byte, packet []byte, h *header.H, fwPacket *firewall.Packet, parsedRx *batch.RxParsed, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
// The entire body is sent as AD, not encrypted. // The entire body is sent as AD, not encrypted.
// The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value. // The packet consists of a 16-byte parsed Nebula header, Associated Data-protected payload, and a trailing 16-byte AEAD signature value.
// The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's // The packet is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
@@ -182,13 +174,6 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
if err != nil { if err != nil {
return 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. // Successfully validated the thing. Get rid of the Relay header.
signedPayload = signedPayload[header.Len:] signedPayload = signedPayload[header.Len:]
// Pull the Roaming parts up here, and return in all call paths. // Pull the Roaming parts up here, and return in all call paths.
@@ -202,7 +187,8 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
// The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing // The only way this happens is if hostmap has an index to the correct HostInfo, but the HostInfo is missing
// its internal mapping. This should never happen. // its internal mapping. This should never happen.
hostinfo.logger(f.l).Error("HostInfo missing remote relay index", hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
"relayRemoteIndex", h.RemoteIndex, "vpnAddrs", hostinfo.vpnAddrs,
"remoteIndex", h.RemoteIndex,
) )
return return
} }
@@ -218,15 +204,16 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
relay: relay, relay: relay,
IsRelayed: true, IsRelayed: true,
} }
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache, meta) f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, parsedRx, lhf, nb, q, localCache, meta)
return
case ForwardingType: case ForwardingType:
// Find the target HostInfo relay object // Find the target HostInfo relay object
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr) targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
if err != nil { if err != nil {
hostinfo.logger(f.l).Info("Failed to find target host info by ip", hostinfo.logger(f.l).Info("Failed to find target host info by ip",
"relayTo", relay.PeerAddr, "relayTo", relay.PeerAddr,
"relayFrom", hostinfo.vpnAddrs[0],
"error", err, "error", err,
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
) )
return return
} }
@@ -277,8 +264,7 @@ func (f *Interface) sendCloseTunnel(h *HostInfo) {
} }
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) { func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
curRemote := hostinfo.GetRemote() if !via.IsRelayed && hostinfo.remote != via.UdpAddr {
if !via.IsRelayed && curRemote != via.UdpAddr {
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) { if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr) hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr)
@@ -290,7 +276,7 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("Suppressing roam back to previous remote", hostinfo.logger(f.l).Debug("Suppressing roam back to previous remote",
"suppressSeconds", RoamingSuppressSeconds, "suppressSeconds", RoamingSuppressSeconds,
"udpAddr", curRemote, "udpAddr", hostinfo.remote,
"newAddr", via.UdpAddr, "newAddr", via.UdpAddr,
) )
} }
@@ -298,11 +284,11 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
} }
hostinfo.logger(f.l).Info("Host roamed to new udp ip/port.", hostinfo.logger(f.l).Info("Host roamed to new udp ip/port.",
"udpAddr", curRemote, "udpAddr", hostinfo.remote,
"newAddr", via.UdpAddr, "newAddr", via.UdpAddr,
) )
hostinfo.lastRoam = time.Now() hostinfo.lastRoam = time.Now()
hostinfo.lastRoamRemote = curRemote hostinfo.lastRoamRemote = hostinfo.remote
hostinfo.SetRemote(via.UdpAddr) hostinfo.SetRemote(via.UdpAddr)
} }
@@ -318,189 +304,16 @@ var (
) )
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers // newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
// newPacket parses data into a fully-hydrated firewall.Packet — kept as a
// thin wrapper around newPacketKey + Hydrate so there's one source of
// parse logic. Callers that don't need the netip.Addr-rich form (e.g.
// conntrack-only paths) should use newPacketKey directly.
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error { func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
if len(data) < 1 { var parsed batch.RxParsed
return ErrPacketTooShort if err := batch.ParsePacket(data, incoming, &parsed); err != nil {
return err
} }
parsed.Key.Hydrate(fp)
version := int((data[0] >> 4) & 0x0f)
switch version {
case ipv4.Version:
return parseV4(data, incoming, fp)
case ipv6.Version:
return parseV6(data, incoming, fp)
}
return ErrUnknownIPVersion
}
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
dataLen := len(data)
if dataLen < ipv6.HeaderLen {
return ErrIPv6PacketTooShort
}
if incoming {
fp.RemoteAddr, _ = netip.AddrFromSlice(data[8:24])
fp.LocalAddr, _ = netip.AddrFromSlice(data[24:40])
} else {
fp.LocalAddr, _ = netip.AddrFromSlice(data[8:24])
fp.RemoteAddr, _ = netip.AddrFromSlice(data[24:40])
}
protoAt := 6 // NextHeader is at 6 bytes into the ipv6 header
offset := ipv6.HeaderLen // Start at the end of the ipv6 header
next := 0
for {
if protoAt >= dataLen {
break
}
proto := layers.IPProtocol(data[protoAt])
switch proto {
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
fp.Protocol = uint8(proto)
fp.RemotePort = 0
fp.LocalPort = 0
fp.Fragment = false
return nil
case layers.IPProtocolICMPv6:
if dataLen < offset+6 {
return ErrIPv6PacketTooShort
}
fp.Protocol = uint8(proto)
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
icmptype := data[offset+1]
switch icmptype {
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply:
fp.RemotePort = binary.BigEndian.Uint16(data[offset+4 : offset+6]) //identifier
default:
fp.RemotePort = 0
}
fp.Fragment = false
return nil
case layers.IPProtocolTCP, layers.IPProtocolUDP:
if dataLen < offset+4 {
return ErrIPv6PacketTooShort
}
fp.Protocol = uint8(proto)
if incoming {
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
} else {
fp.LocalPort = binary.BigEndian.Uint16(data[offset : offset+2])
fp.RemotePort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
}
fp.Fragment = false
return nil
case layers.IPProtocolIPv6Fragment:
// Fragment header is 8 bytes, need at least offset+4 to read the offset field
if dataLen < offset+8 {
return ErrIPv6PacketTooShort
}
// Check if this is the first fragment
fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
if fragmentOffset != 0 {
// Non-first fragment, use what we have now and stop processing
fp.Protocol = data[offset]
fp.Fragment = true
fp.RemotePort = 0
fp.LocalPort = 0
return nil
}
// The next loop should be the transport layer since we are the first fragment
next = 8 // Fragment headers are always 8 bytes
case layers.IPProtocolAH:
// Auth headers, used by IPSec, have a different meaning for header length
if dataLen <= offset+1 {
break
}
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
}
if next <= 0 {
// Safety check, each ipv6 header has to be at least 8 bytes
next = 8
}
protoAt = offset
offset = offset + next
}
return ErrIPv6CouldNotFindPayload
}
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
// Do we at least have an ipv4 header worth of data?
if len(data) < ipv4.HeaderLen {
return ErrIPv4PacketTooShort
}
// Adjust our start position based on the advertised ip header length
ihl := int(data[0]&0x0f) << 2
// Well-formed ip header length?
if ihl < ipv4.HeaderLen {
return ErrIPv4InvalidHeaderLength
}
// Check if this is the second or further fragment of a fragmented packet.
flagsfrags := binary.BigEndian.Uint16(data[6:8])
fp.Fragment = (flagsfrags & 0x1FFF) != 0
// Firewall handles protocol checks
fp.Protocol = data[9]
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
minLen := ihl
if !fp.Fragment {
if fp.Protocol == firewall.ProtoICMP {
minLen += minFwPacketLen + 2
} else {
minLen += minFwPacketLen
}
}
if len(data) < minLen {
return ErrIPv4InvalidHeaderLength
}
if incoming { // Firewall packets are locally oriented
fp.RemoteAddr, _ = netip.AddrFromSlice(data[12:16])
fp.LocalAddr, _ = netip.AddrFromSlice(data[16:20])
} else {
fp.LocalAddr, _ = netip.AddrFromSlice(data[12:16])
fp.RemoteAddr, _ = netip.AddrFromSlice(data[16:20])
}
if fp.Fragment {
fp.RemotePort = 0
fp.LocalPort = 0
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
} else if incoming {
fp.RemotePort = binary.BigEndian.Uint16(data[ihl : ihl+2]) //src port
fp.LocalPort = binary.BigEndian.Uint16(data[ihl+2 : ihl+4]) //dst port
} else {
fp.LocalPort = binary.BigEndian.Uint16(data[ihl : ihl+2]) //src port
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+2 : ihl+4]) //dst port
}
return nil return nil
} }
@@ -549,23 +362,7 @@ func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger
case ecnCE: case ecnCE:
// Already CE. // Already CE.
default: default:
// Rewriting the ToS byte invalidates the IPv4 header checksum, so
// patch it incrementally per RFC 1624 (HC' = ~(~HC + ~m + m')). The
// ToS is the low byte of the 16-bit word at pkt[0:2]; the header
// checksum lives at pkt[10:12]. A header too short to carry a
// checksum can't be fixed up here, so leave it for newPacket to
// reject rather than emit a mangled packet.
if len(pkt) < ipv4.HeaderLen {
return
}
m := binary.BigEndian.Uint16(pkt[0:2])
pkt[1] = (pkt[1] &^ 0x03) | ecnCE pkt[1] = (pkt[1] &^ 0x03) | ecnCE
mNew := binary.BigEndian.Uint16(pkt[0:2])
sum := uint32(^binary.BigEndian.Uint16(pkt[10:12])) + uint32(^m) + uint32(mNew)
for sum > 0xffff {
sum = (sum >> 16) + (sum & 0xffff)
}
binary.BigEndian.PutUint16(pkt[10:12], ^uint16(sum))
} }
case 6: case 6:
switch (pkt[1] >> 4) & 0x03 { switch (pkt[1] >> 4) & 0x03 {
@@ -581,7 +378,7 @@ func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger
} }
} }
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) { func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, parsedRx *batch.RxParsed, nb []byte, q int, localCache firewall.ConntrackCache, meta udp.RxMeta) {
// RFC 6040 normal-mode combine: fold any outer CE mark stamped by the // RFC 6040 normal-mode combine: fold any outer CE mark stamped by the
// underlay into the inner header before firewall + TUN write. Other // underlay into the inner header before firewall + TUN write. Other
// outer codepoints are advisory only — we keep the inner unchanged. // outer codepoints are advisory only — we keep the inner unchanged.
@@ -589,7 +386,13 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
applyOuterECN(out, meta.OuterECN, hostinfo, f.l) applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
} }
err := newPacket(out, true, fwPacket) // Single IP+L4 walk feeds the firewall conntrack key (parsedRx.Key)
// and the batcher hint (parsedRx.tcp/udp). Replaces newPacket — and
// pointedly does NOT fill fwPacket.LocalAddr/RemoteAddr, since
// firewall.Drop's fast path uses Key alone and only hydrates fwPacket
// from Key on the slow path.
*fwPacket = firewall.Packet{}
err := batch.ParsePacket(out, true, parsedRx)
if err != nil { if err != nil {
hostinfo.logger(f.l).Warn("Error while validating inbound packet", hostinfo.logger(f.l).Warn("Error while validating inbound packet",
"error", err, "error", err,
@@ -598,13 +401,11 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
return return
} }
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache) dropReason := f.firewall.Drop(parsedRx.Key, fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
if dropReason != nil { if dropReason != nil {
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore // NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
// This gives us a buffer to build the reject packet in. With UDP GRO this is a single segment of a shared // This gives us a buffer to build the reject packet in
// recvmmsg row whose capacity runs to the end of the whole row, so cap it to its own length (cap==len) to f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
// keep the reject builder from writing past this segment into the next, not-yet-processed coalesced segment.
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet[:len(packet):len(packet)], q)
if f.l.Enabled(context.Background(), slog.LevelDebug) { if f.l.Enabled(context.Background(), slog.LevelDebug) {
hostinfo.logger(f.l).Debug("dropping inbound packet", hostinfo.logger(f.l).Debug("dropping inbound packet",
"fwPacket", fwPacket, "fwPacket", fwPacket,
@@ -614,7 +415,7 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
return return
} }
err = f.batchers[q].Commit(out) err = f.batchers[q].CommitInbound(out, parsedRx)
if err != nil { if err != nil {
f.l.Error("Failed to write to tun", "error", err) f.l.Error("Failed to write to tun", "error", err)
} }
@@ -661,11 +462,10 @@ func (f *Interface) handleRecvError(addr netip.AddrPort, h *header.H) {
return return
} }
hr := hostinfo.GetRemote() if hostinfo.remote.IsValid() && hostinfo.remote != addr {
if hr.IsValid() && hr != addr {
f.l.Info("Someone spoofing recv_errors?", f.l.Info("Someone spoofing recv_errors?",
"addr", addr, "addr", addr,
"hostinfoRemote", hr, "hostinfoRemote", hostinfo.remote,
) )
return return
} }
+20 -54
View File
@@ -11,6 +11,7 @@ import (
"github.com/google/gopacket/layers" "github.com/google/gopacket/layers"
"github.com/slackhq/nebula/firewall" "github.com/slackhq/nebula/firewall"
"github.com/slackhq/nebula/overlay/batch"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"golang.org/x/net/ipv4" "golang.org/x/net/ipv4"
@@ -21,13 +22,13 @@ func Test_newPacket(t *testing.T) {
// length fails // length fails
err := newPacket([]byte{}, true, p) err := newPacket([]byte{}, true, p)
require.ErrorIs(t, err, ErrPacketTooShort) require.ErrorIs(t, err, batch.ErrPacketTooShort)
err = newPacket([]byte{0x40}, true, p) err = newPacket([]byte{0x40}, true, p)
require.ErrorIs(t, err, ErrIPv4PacketTooShort) require.ErrorIs(t, err, batch.ErrIPv4PacketTooShort)
err = newPacket([]byte{0x60}, true, p) err = newPacket([]byte{0x60}, true, p)
require.ErrorIs(t, err, ErrIPv6PacketTooShort) require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
// length fail with ip options // length fail with ip options
h := ipv4.Header{ h := ipv4.Header{
@@ -40,15 +41,15 @@ func Test_newPacket(t *testing.T) {
b, _ := h.Marshal() b, _ := h.Marshal()
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength) require.ErrorIs(t, err, batch.ErrIPv4InvalidHeaderLength)
// not an ipv4 packet // not an ipv4 packet
err = newPacket([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p) err = newPacket([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p)
require.ErrorIs(t, err, ErrUnknownIPVersion) require.ErrorIs(t, err, batch.ErrUnknownIPVersion)
// invalid ihl // invalid ihl
err = newPacket([]byte{4<<4 | (8 >> 2 & 0x0f), 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p) err = newPacket([]byte{4<<4 | (8 >> 2 & 0x0f), 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p)
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength) require.ErrorIs(t, err, batch.ErrIPv4InvalidHeaderLength)
// account for variable ip header length - incoming // account for variable ip header length - incoming
h = ipv4.Header{ h = ipv4.Header{
@@ -115,7 +116,7 @@ func Test_newPacket_v6(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
err = newPacket(buffer.Bytes(), true, p) err = newPacket(buffer.Bytes(), true, p)
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
// A v6 packet with a hop-by-hop extension // A v6 packet with a hop-by-hop extension
// ICMPv6 Payload (Echo Request) // ICMPv6 Payload (Echo Request)
@@ -149,12 +150,12 @@ func Test_newPacket_v6(t *testing.T) {
// A full IPv6 header and 1 byte in the first extension, but missing // A full IPv6 header and 1 byte in the first extension, but missing
// the length byte. // the length byte.
err = newPacket(buffer.Bytes()[:41], true, p) err = newPacket(buffer.Bytes()[:41], true, p)
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
// A full IPv6 header plus 1 full extension, but only 1 byte of the // A full IPv6 header plus 1 full extension, but only 1 byte of the
// next layer, missing length byte // next layer, missing length byte
err = newPacket(buffer.Bytes()[:49], true, p) err = newPacket(buffer.Bytes()[:49], true, p)
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
err = nil err = nil
// A good ICMP packet // A good ICMP packet
@@ -217,7 +218,7 @@ func Test_newPacket_v6(t *testing.T) {
b = buffer.Bytes() b = buffer.Bytes()
b[6] = 255 // 255 is a reserved protocol number b[6] = 255 // 255 is a reserved protocol number
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
// A good UDP packet // A good UDP packet
ip = layers.IPv6{ ip = layers.IPv6{
@@ -264,7 +265,7 @@ func Test_newPacket_v6(t *testing.T) {
// Too short UDP packet // Too short UDP packet
err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes
require.ErrorIs(t, err, ErrIPv6PacketTooShort) require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
// A good TCP packet // A good TCP packet
b[6] = byte(layers.IPProtocolTCP) b[6] = byte(layers.IPProtocolTCP)
@@ -291,7 +292,7 @@ func Test_newPacket_v6(t *testing.T) {
// Too short TCP packet // Too short TCP packet
err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes
require.ErrorIs(t, err, ErrIPv6PacketTooShort) require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
// A good UDP packet with an AH header // A good UDP packet with an AH header
ip = layers.IPv6{ ip = layers.IPv6{
@@ -336,12 +337,12 @@ func Test_newPacket_v6(t *testing.T) {
// Ensure buffer bounds checking during processing // Ensure buffer bounds checking during processing
err = newPacket(b[:41], true, p) err = newPacket(b[:41], true, p)
require.ErrorIs(t, err, ErrIPv6PacketTooShort) require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
// Invalid AH header // Invalid AH header
b = buffer.Bytes() b = buffer.Bytes()
err = newPacket(b, true, p) err = newPacket(b, true, p)
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload) require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
} }
func Test_newPacket_ipv6Fragment(t *testing.T) { func Test_newPacket_ipv6Fragment(t *testing.T) {
@@ -448,7 +449,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) {
// Too short of a fragment packet // Too short of a fragment packet
err = newPacket(secondFrag[:len(secondFrag)-10], false, p) err = newPacket(secondFrag[:len(secondFrag)-10], false, p)
require.ErrorIs(t, err, ErrIPv6PacketTooShort) require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
} }
func BenchmarkParseV6(b *testing.B) { func BenchmarkParseV6(b *testing.B) {
@@ -529,7 +530,7 @@ func BenchmarkParseV6(b *testing.B) {
b.Run("Normal", func(b *testing.B) { b.Run("Normal", func(b *testing.B) {
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
if err = parseV6(normalPacket, true, fp); err != nil { if err = newPacket(normalPacket, true, fp); err != nil {
b.Fatal(err) b.Fatal(err)
} }
} }
@@ -537,7 +538,7 @@ func BenchmarkParseV6(b *testing.B) {
b.Run("FirstFragment", func(b *testing.B) { b.Run("FirstFragment", func(b *testing.B) {
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
if err = parseV6(firstFrag, true, fp); err != nil { if err = newPacket(firstFrag, true, fp); err != nil {
b.Fatal(err) b.Fatal(err)
} }
} }
@@ -545,7 +546,7 @@ func BenchmarkParseV6(b *testing.B) {
b.Run("SecondFragment", func(b *testing.B) { b.Run("SecondFragment", func(b *testing.B) {
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
if err = parseV6(secondFrag, true, fp); err != nil { if err = newPacket(secondFrag, true, fp); err != nil {
b.Fatal(err) b.Fatal(err)
} }
} }
@@ -590,7 +591,7 @@ func BenchmarkParseV6(b *testing.B) {
b.Run("200 HopByHop headers", func(b *testing.B) { b.Run("200 HopByHop headers", func(b *testing.B) {
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
if err = parseV6(evilBytes, false, fp); err != nil { if err = newPacket(evilBytes, false, fp); err != nil {
b.Fatal(err) b.Fatal(err)
} }
} }
@@ -640,38 +641,3 @@ func serializeAH(ah *layers.IPSecAH) []byte {
return buf.Bytes() 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")
}
+14 -7
View File
@@ -5,11 +5,18 @@ import "net/netip"
type RxBatcher interface { type RxBatcher interface {
// Reserve creates a pkt to borrow // Reserve creates a pkt to borrow
Reserve(sz int) []byte Reserve(sz int) []byte
// Commit borrows pkt. The caller must keep pkt valid until the next Flush // Commit borrows pkt. The caller must keep pkt valid until the next Flush.
// Walks IP+L4 headers itself; prefer CommitInbound when the caller already
// has an RxParsed in hand from ParsePacket.
Commit(pkt []byte) error Commit(pkt []byte) error
// Flush emits every queued packet in arrival order. // CommitInbound is Commit with a hint produced by ParsePacket, so the
// Returns the first error observed; keeps draining so one bad packet doesn't hold up the rest. // batcher can skip the IP+L4 re-parse. Borrowed slice contract is the
// After Flush returns, borrowed payload slices may be recycled. // same as Commit. Implementations that don't coalesce may delegate to
// Commit.
CommitInbound(pkt []byte, parsed *RxParsed) error
// Flush emits every queued packet in arrival order. Returns the
// first error observed; keeps draining so one bad packet doesn't hold up
// the rest. After Flush returns, borrowed payload slices may be recycled.
Flush() error Flush() error
} }
@@ -21,8 +28,8 @@ type TxBatcher interface {
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT) // caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
// to leave the outer ECN field unset. // to leave the outer ECN field unset.
Commit(pkt []byte, dst netip.AddrPort, outerECN byte) Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
// Flush emits every queued packet via the underlying batch writer in arrival order. // Flush emits every queued packet via the underlying batch writer in
// Returns an errors.Join of one or more errors. // arrival order. Returns an errors.Join of one or more errors. After Flush returns,
// After Flush returns, borrowed payload slices may be recycled. // borrowed payload slices may be recycled.
Flush() error Flush() error
} }
+27 -45
View File
@@ -93,11 +93,8 @@ func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
// ipHeadersMatch compares the IP portion of two packet header prefixes for // ipHeadersMatch compares the IP portion of two packet header prefixes for
// byte-for-byte equality on every field that must be identical across // byte-for-byte equality on every field that must be identical across
// coalesced segments. Size/IPID/IPCsum are masked out. The full DSCP/ECN // coalesced segments. Size/IPID/IPCsum and the 2-bit IP-level ECN field are
// byte (IPv4 ToS / IPv6 traffic class) is compared, matching Linux kernel // masked out — the appendPayload step merges CE into the seed.
// GRO: segments with differing ECN codepoints must not coalesce, otherwise
// ORing e.g. ECT(0) with ECT(1) would fabricate a false CE (congestion)
// mark or mark a Not-ECT flow as ECN-capable.
// //
// The transport (L4) portion of the header is checked separately by the // The transport (L4) portion of the header is checked separately by the
// per-protocol matcher. // per-protocol matcher.
@@ -105,11 +102,11 @@ func ipHeadersMatch(a, b []byte, isV6 bool) bool {
if isV6 { if isV6 {
// IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16], // IPv6: byte 0 = version/TC[7:4], byte 1 = TC[3:0]/flow[19:16],
// bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst. // bytes [2:4] = flow[15:0], [6:8] = next_hdr/hop, [8:40] = src+dst.
// Compare byte 1 fully so ECN (TC[1:0]) must match. Skip [4:6] payload_len. // ECN lives in TC[1:0] = byte 1 mask 0x30. Skip [4:6] payload_len.
if a[0] != b[0] { if a[0] != b[0] {
return false return false
} }
if a[1] != b[1] { if a[1]&^0x30 != b[1]&^0x30 {
return false return false
} }
if !bytes.Equal(a[2:4], b[2:4]) { if !bytes.Equal(a[2:4], b[2:4]) {
@@ -122,12 +119,11 @@ func ipHeadersMatch(a, b []byte, isV6 bool) bool {
} }
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2), // IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst. // [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
// Compare byte 1 fully so ECN must match.
// Skip [2:4] total len, [4:6] id, [10:12] csum. // Skip [2:4] total len, [4:6] id, [10:12] csum.
if a[0] != b[0] { if a[0] != b[0] {
return false return false
} }
if a[1] != b[1] { if a[1]&^0x03 != b[1]&^0x03 {
return false return false
} }
if !bytes.Equal(a[6:10], b[6:10]) { if !bytes.Equal(a[6:10], b[6:10]) {
@@ -139,43 +135,29 @@ func ipHeadersMatch(a, b []byte, isV6 bool) bool {
return true return true
} }
// Arena is an injectable byte-slab that hands out non-overlapping borrowed // mergeECNIntoSeed ORs the 2-bit IP-level ECN field of pkt's IP header
// slices via Reserve and releases them in bulk via Reset. // onto the seed's IP header, so a CE mark on any coalesced segment
type Arena struct { // propagates to the final superpacket. (CE is 0b11; ORing yields CE if
buf []byte // any segment carried it.) Used by both TCP and UDP coalescers, so the
// invariant lives in one place.
func mergeECNIntoSeed(seedHdr, pktHdr []byte, isV6 bool) {
if isV6 {
seedHdr[1] |= pktHdr[1] & 0x30
} else {
seedHdr[1] |= pktHdr[1] & 0x03
}
} }
// NewArena returns an Arena with a pre-allocated backing of the given // reserveFromBacking implements the Reserve half of the RxBatcher contract
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that // shared by TCP and UDP coalescers. The backing slice grows on demand;
// only feeds the coalescer pre-made []byte packets via Commit). // already-committed slices reference the old array and remain valid until
func NewArena(capacity int) *Arena { // Flush resets backing.
return &Arena{buf: make([]byte, 0, capacity)} func reserveFromBacking(backing *[]byte, sz int) []byte {
if len(*backing)+sz > cap(*backing) {
newCap := max(cap(*backing)*2, sz)
*backing = make([]byte, 0, newCap)
} }
start := len(*backing)
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the *backing = (*backing)[:start+sz]
// request doesn't fit the current backing, a fresh, larger backing is return (*backing)[start : start+sz : start+sz]
// allocated; already-borrowed slices reference the old backing and remain
// valid until Reset.
func (a *Arena) Reserve(sz int) []byte {
if len(a.buf)+sz > cap(a.buf) {
newCap := max(cap(a.buf)*2, sz)
a.buf = make([]byte, 0, newCap)
} }
start := len(a.buf)
a.buf = a.buf[:start+sz]
return a.buf[start : start+sz : start+sz]
}
// Reset releases every slice handed out since the last Reset. Callers must
// not use any previously-borrowed slice after this returns. The underlying
// backing array is retained so subsequent Reserves don't re-allocate.
func (a *Arena) Reset() {
a.buf = a.buf[:0]
}
// Reserver hands out an sz-byte slice valid until its Resetter runs.
type Reserver func(sz int) []byte
// Resetter clears all reservations held by a Reserver. Only the arena's
// owner holds one; lanes inside a MultiCoalescer get nil.
type Resetter func()
+443
View File
@@ -0,0 +1,443 @@
package batch
import (
"encoding/binary"
"errors"
"github.com/slackhq/nebula/firewall"
)
// IANA protocol numbers we recognise during the inbound parse. Kept local
// (rather than reaching for the firewall constants for every one of these)
// so the byte-comparison hot path doesn't depend on cross-package values.
const (
ipProtoICMP = 1
ipProtoIPv6Fragment = 44
ipProtoESP = 50
ipProtoAH = 51
ipProtoICMPv6 = 58
ipProtoNoNextHdr = 59
icmpv6TypeEchoRequest = 128
icmpv6TypeEchoReply = 129
)
// Packet parse errors — the canonical sentinel set for IP+L4 parsing.
// Both inbound and outbound callers share this surface, so any code path
// that ends up at firewall.PacketKey reports drops with the same errors.
var (
ErrPacketTooShort = errors.New("packet is too short")
ErrUnknownIPVersion = errors.New("packet is an unknown ip version")
ErrIPv4InvalidHeaderLength = errors.New("invalid ipv4 header length")
ErrIPv4PacketTooShort = errors.New("ipv4 packet is too short")
ErrIPv6PacketTooShort = errors.New("ipv6 packet is too short")
ErrIPv6CouldNotFindPayload = errors.New("could not find payload in ipv6 packet")
)
// RxKind discriminates how an inbound plaintext packet should be committed
// after its firewall.Packet has been built. RxKindPassthrough means the
// IP shape is valid (firewall could match on it) but the coalescer's
// strict checks reject it — caller should still write it via the
// passthrough lane.
type RxKind uint8
const (
RxKindPassthrough RxKind = iota
RxKindTCP
RxKindUDP
)
// RxParsed is the unified result of one IP+L4 walk:
// - Key: the firewall's conntrack/cache lookup key. The dense form lets
// firewall.Drop hit conntrack without ever filling the rich Packet's
// netip.Addr fields. On a conntrack miss, Drop hydrates the caller's
// Packet from Key.
// - tcp/udp: the coalescer hint so commitParsed doesn't re-walk the
// headers. Meaningful only when Kind is RxKindTCP / RxKindUDP.
type RxParsed struct {
Kind RxKind
Key firewall.PacketKey
tcp parsedTCP
udp parsedUDP
}
// ParsePacket walks an IP packet once and fills parsed.Key. When incoming
// is true and the L4 shape is coalesce-eligible, also fills parsed.tcp /
// parsed.udp so CommitInbound can dispatch into the coalescer without
// re-walking the headers.
//
// Direction selects the Key orientation:
//
// incoming=true → wire src → Key.RemoteAddr/Port, wire dst → Key.LocalAddr/Port
// incoming=false → wire src → Key.LocalAddr/Port, wire dst → Key.RemoteAddr/Port
//
// ICMP always lands the identifier in Key.RemotePort, regardless of direction.
//
// Eligibility rules for the coalescer hint match the coalescer's own
// parseTCPBase/parseUDP:
// - IPv4 strict: IHL == 20, no fragmentation (MF or offset), proto TCP/UDP.
// - IPv6 strict: NextHeader is directly TCP or UDP (no extension headers).
//
// The hint is only filled for incoming packets, since the outbound path
// does not feed an inbound coalescer. Outbound callers see Kind stay at
// RxKindPassthrough and parsed.tcp/udp stay zero.
func ParsePacket(pkt []byte, incoming bool, parsed *RxParsed) error {
parsed.Kind = RxKindPassthrough
// Reset Key in full: v4 only writes the low 4 bytes of each address
// field, so without this a v6 call followed by a v4 reusing the same
// RxParsed would inherit the high 12 bytes — breaking the conntrack
// map equality for v4 flows.
parsed.Key = firewall.PacketKey{}
if len(pkt) < 1 {
return ErrPacketTooShort
}
switch pkt[0] >> 4 {
case 4:
return parsePacketV4(pkt, incoming, parsed)
case 6:
return parsePacketV6(pkt, incoming, parsed)
}
return ErrUnknownIPVersion
}
// parsePacketV4 fills parsed.Key from an IPv4 packet. Direction selects
// Local/Remote orientation. When incoming and the shape is strict, also
// fills the coalescer hint.
func parsePacketV4(pkt []byte, incoming bool, parsed *RxParsed) error {
if len(pkt) < 20 {
return ErrIPv4PacketTooShort
}
ihl := int(pkt[0]&0x0f) << 2
if ihl < 20 {
return ErrIPv4InvalidHeaderLength
}
flagsfrags := binary.BigEndian.Uint16(pkt[6:8])
parsed.Key.Fragment = (flagsfrags & 0x1FFF) != 0
parsed.Key.Protocol = pkt[9]
parsed.Key.IsV6 = false
// minFwPacketLen (4) is the L4-header prefix the firewall needs to pull
// ports; ICMP needs two extra bytes for the identifier.
minLen := ihl
if !parsed.Key.Fragment {
if parsed.Key.Protocol == firewall.ProtoICMP {
minLen += 4 + 2
} else {
minLen += 4
}
}
if len(pkt) < minLen {
return ErrIPv4InvalidHeaderLength
}
if incoming {
copy(parsed.Key.RemoteAddr[:4], pkt[12:16])
copy(parsed.Key.LocalAddr[:4], pkt[16:20])
} else {
copy(parsed.Key.LocalAddr[:4], pkt[12:16])
copy(parsed.Key.RemoteAddr[:4], pkt[16:20])
}
switch {
case parsed.Key.Fragment:
parsed.Key.RemotePort = 0
parsed.Key.LocalPort = 0
case parsed.Key.Protocol == firewall.ProtoICMP:
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[ihl+4 : ihl+6])
parsed.Key.LocalPort = 0
case incoming:
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[ihl : ihl+2])
parsed.Key.LocalPort = binary.BigEndian.Uint16(pkt[ihl+2 : ihl+4])
default:
parsed.Key.LocalPort = binary.BigEndian.Uint16(pkt[ihl : ihl+2])
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[ihl+2 : ihl+4])
}
// Coalescer hint is inbound-only: no inbound coalescer fires on outgoing.
if !incoming {
return nil
}
// Coalescer-eligible? Strict shape: IHL==20, no MF/offset, TCP or UDP.
if ihl != 20 || (flagsfrags&0x3FFF) != 0 {
return nil
}
if parsed.Key.Protocol != ipProtoTCP && parsed.Key.Protocol != ipProtoUDP {
return nil
}
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
if totalLen > len(pkt) || totalLen < 20 {
return nil
}
pktTrim := pkt[:totalLen]
switch parsed.Key.Protocol {
case ipProtoTCP:
fillParsedTCPv4(pktTrim, parsed)
case ipProtoUDP:
fillParsedUDPv4(pktTrim, parsed)
}
return nil
}
// fillParsedTCPv4 fills parsed.tcp from a strict-shape IPv4+TCP packet
// already validated to have IHL==20 and to be totalLen-trimmed.
func fillParsedTCPv4(pkt []byte, parsed *RxParsed) {
if len(pkt) < 40 { // IPv4(20) + min TCP(20)
return
}
tcpOff := int(pkt[32]>>4) * 4
if tcpOff < 20 || tcpOff > 60 {
return
}
if len(pkt) < 20+tcpOff {
return
}
p := &parsed.tcp
p.ipHdrLen = 20
p.tcpHdrLen = tcpOff
p.hdrLen = 20 + tcpOff
p.payLen = len(pkt) - p.hdrLen
p.seq = binary.BigEndian.Uint32(pkt[24:28])
p.flags = pkt[33]
p.fk.isV6 = false
p.fk.sport = parsed.Key.RemotePort
p.fk.dport = parsed.Key.LocalPort
copy(p.fk.src[:4], pkt[12:16])
copy(p.fk.dst[:4], pkt[16:20])
parsed.Kind = RxKindTCP
}
// fillParsedUDPv4 fills parsed.udp from a strict-shape IPv4+UDP packet.
func fillParsedUDPv4(pkt []byte, parsed *RxParsed) {
if len(pkt) < 28 { // IPv4(20) + UDP(8)
return
}
udpLen := int(binary.BigEndian.Uint16(pkt[24:26]))
if udpLen < 8 || udpLen > len(pkt)-20 {
return
}
p := &parsed.udp
p.ipHdrLen = 20
p.hdrLen = 28
p.payLen = udpLen - 8
p.fk.isV6 = false
p.fk.sport = parsed.Key.RemotePort
p.fk.dport = parsed.Key.LocalPort
copy(p.fk.src[:4], pkt[12:16])
copy(p.fk.dst[:4], pkt[16:20])
parsed.Kind = RxKindUDP
}
// parsePacketV6 fills parsed.Key from an IPv6 packet. Direction selects
// Local/Remote orientation. The coalescer hint fast path only triggers
// when NextHeader is directly TCP or UDP — any extension header chain
// falls into the lenient walk below, and the hint stays unfilled.
func parsePacketV6(pkt []byte, incoming bool, parsed *RxParsed) error {
if len(pkt) < 40 {
return ErrIPv6PacketTooShort
}
parsed.Key.IsV6 = true
if incoming {
copy(parsed.Key.RemoteAddr[:], pkt[8:24])
copy(parsed.Key.LocalAddr[:], pkt[24:40])
} else {
copy(parsed.Key.LocalAddr[:], pkt[8:24])
copy(parsed.Key.RemoteAddr[:], pkt[24:40])
}
if proto := pkt[6]; proto == ipProtoTCP || proto == ipProtoUDP {
// Strict v6: ports are at the IP header end. Always fill key; only
// fill the coalescer hint if the L4 shape passes.
if len(pkt) < 44 {
return ErrIPv6PacketTooShort
}
parsed.Key.Protocol = proto
parsed.Key.Fragment = false
if incoming {
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[40:42])
parsed.Key.LocalPort = binary.BigEndian.Uint16(pkt[42:44])
} else {
parsed.Key.LocalPort = binary.BigEndian.Uint16(pkt[40:42])
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[42:44])
}
// Coalescer hint is inbound-only.
if !incoming {
return nil
}
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
if 40+payloadLen > len(pkt) {
return nil
}
pktTrim := pkt[:40+payloadLen]
switch proto {
case ipProtoTCP:
fillParsedTCPv6(pktTrim, parsed)
case ipProtoUDP:
fillParsedUDPv6(pktTrim, parsed)
}
return nil
}
// Slow path: walk extension header chain. Coalescer hint never fires
// here, so direction only matters for L4 port orientation.
return walkV6Headers(pkt, incoming, parsed)
}
func fillParsedTCPv6(pkt []byte, parsed *RxParsed) {
if len(pkt) < 60 { // IPv6(40) + min TCP(20)
return
}
tcpOff := int(pkt[52]>>4) * 4
if tcpOff < 20 || tcpOff > 60 {
return
}
if len(pkt) < 40+tcpOff {
return
}
p := &parsed.tcp
p.ipHdrLen = 40
p.tcpHdrLen = tcpOff
p.hdrLen = 40 + tcpOff
p.payLen = len(pkt) - p.hdrLen
p.seq = binary.BigEndian.Uint32(pkt[44:48])
p.flags = pkt[53]
p.fk.isV6 = true
p.fk.sport = parsed.Key.RemotePort
p.fk.dport = parsed.Key.LocalPort
copy(p.fk.src[:], pkt[8:24])
copy(p.fk.dst[:], pkt[24:40])
parsed.Kind = RxKindTCP
}
func fillParsedUDPv6(pkt []byte, parsed *RxParsed) {
if len(pkt) < 48 { // IPv6(40) + UDP(8)
return
}
udpLen := int(binary.BigEndian.Uint16(pkt[44:46]))
if udpLen < 8 || udpLen > len(pkt)-40 {
return
}
p := &parsed.udp
p.ipHdrLen = 40
p.hdrLen = 48
p.payLen = udpLen - 8
p.fk.isV6 = true
p.fk.sport = parsed.Key.RemotePort
p.fk.dport = parsed.Key.LocalPort
copy(p.fk.src[:], pkt[8:24])
copy(p.fk.dst[:], pkt[24:40])
parsed.Kind = RxKindUDP
}
// walkV6Headers handles every IPv6 case the strict "NextHeader == TCP/UDP"
// fast path doesn't: ESP, NoNextHeader, ICMPv6, fragment headers (first vs
// later), AH, generic extension headers. Coalescer eligibility is always
// RxKindPassthrough on this path (parsed already initialised that way).
// Direction matters only for the L4 port orientation when the chain
// terminates at TCP/UDP.
func walkV6Headers(pkt []byte, incoming bool, parsed *RxParsed) error {
dataLen := len(pkt)
protoAt := 6
offset := 40
next := 0
for {
if protoAt >= dataLen {
break
}
proto := pkt[protoAt]
switch proto {
case ipProtoESP, ipProtoNoNextHdr:
parsed.Key.Protocol = proto
parsed.Key.RemotePort = 0
parsed.Key.LocalPort = 0
parsed.Key.Fragment = false
return nil
case ipProtoICMPv6:
if dataLen < offset+6 {
return ErrIPv6PacketTooShort
}
parsed.Key.Protocol = proto
parsed.Key.LocalPort = 0
switch pkt[offset+1] {
case icmpv6TypeEchoRequest, icmpv6TypeEchoReply:
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[offset+4 : offset+6])
default:
parsed.Key.RemotePort = 0
}
parsed.Key.Fragment = false
return nil
case ipProtoTCP, ipProtoUDP:
// Reachable when an extension-header chain ends at TCP/UDP. The
// strict-eligible fast path above already handled the no-extension
// case; here we only fill firewall ports and stay passthrough.
if dataLen < offset+4 {
return ErrIPv6PacketTooShort
}
parsed.Key.Protocol = proto
if incoming {
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[offset : offset+2])
parsed.Key.LocalPort = binary.BigEndian.Uint16(pkt[offset+2 : offset+4])
} else {
parsed.Key.LocalPort = binary.BigEndian.Uint16(pkt[offset : offset+2])
parsed.Key.RemotePort = binary.BigEndian.Uint16(pkt[offset+2 : offset+4])
}
parsed.Key.Fragment = false
return nil
case ipProtoIPv6Fragment:
if dataLen < offset+8 {
return ErrIPv6PacketTooShort
}
fragmentOffset := binary.BigEndian.Uint16(pkt[offset+2:offset+4]) &^ uint16(0x7)
if fragmentOffset != 0 {
// Non-first fragment: report the fragment flag and stop.
parsed.Key.Protocol = pkt[offset]
parsed.Key.Fragment = true
parsed.Key.RemotePort = 0
parsed.Key.LocalPort = 0
return nil
}
next = 8
case ipProtoAH:
if dataLen <= offset+1 {
break
}
next = int(pkt[offset+1]+2) << 2
default:
if dataLen <= offset+1 {
break
}
next = int(pkt[offset+1]+1) << 3
}
if next <= 0 {
next = 8
}
protoAt = offset
offset = offset + next
}
return ErrIPv6CouldNotFindPayload
}
// CommitInbound dispatches pkt to the appropriate lane using parsed.Kind,
// skipping the IP+L4 re-parse that MultiCoalescer.Commit would otherwise
// do. Borrowed slice contract is identical to MultiCoalescer.Commit.
func (m *MultiCoalescer) CommitInbound(pkt []byte, parsed *RxParsed) error {
switch parsed.Kind {
case RxKindTCP:
if m.tcp != nil {
return m.tcp.commitParsed(pkt, parsed.tcp)
}
case RxKindUDP:
if m.udp != nil {
return m.udp.commitParsed(pkt, parsed.udp)
}
}
return m.pt.Commit(pkt)
}
+394
View File
@@ -0,0 +1,394 @@
package batch
import (
"encoding/binary"
"net/netip"
"testing"
"github.com/slackhq/nebula/firewall"
)
// parseV4InboundBaseline mirrors what outside.go's parseV4(incoming=true)
// does, so the "split" bench measures the *current* state: firewall-side
// parse, then m.Commit re-parses inside the coalescer. Two walks per
// packet. Kept faithful in shape (one read per field, AddrFromSlice for
// the addrs) so the CPU profile matches the production parseV4.
func parseV4InboundBaseline(pkt []byte, fp *firewall.Packet) bool {
if len(pkt) < 20 {
return false
}
ihl := int(pkt[0]&0x0f) << 2
if ihl < 20 {
return false
}
flagsfrags := binary.BigEndian.Uint16(pkt[6:8])
fp.Fragment = (flagsfrags & 0x1FFF) != 0
fp.Protocol = pkt[9]
minLen := ihl
if !fp.Fragment {
if fp.Protocol == firewall.ProtoICMP {
minLen += 4 + 2
} else {
minLen += 4
}
}
if len(pkt) < minLen {
return false
}
fp.RemoteAddr, _ = netip.AddrFromSlice(pkt[12:16])
fp.LocalAddr, _ = netip.AddrFromSlice(pkt[16:20])
switch {
case fp.Fragment:
fp.RemotePort = 0
fp.LocalPort = 0
case fp.Protocol == firewall.ProtoICMP:
fp.RemotePort = binary.BigEndian.Uint16(pkt[ihl+4 : ihl+6])
fp.LocalPort = 0
default:
fp.RemotePort = binary.BigEndian.Uint16(pkt[ihl : ihl+2])
fp.LocalPort = binary.BigEndian.Uint16(pkt[ihl+2 : ihl+4])
}
return true
}
// parseV6InboundBaseline is the v6 analogue: replicates parseV6's
// extension-header walk so the split bench captures its true cost.
func parseV6InboundBaseline(pkt []byte, fp *firewall.Packet) bool {
dataLen := len(pkt)
if dataLen < 40 {
return false
}
fp.RemoteAddr, _ = netip.AddrFromSlice(pkt[8:24])
fp.LocalAddr, _ = netip.AddrFromSlice(pkt[24:40])
protoAt := 6
offset := 40
next := 0
for {
if protoAt >= dataLen {
return false
}
proto := pkt[protoAt]
switch proto {
case ipProtoESP, ipProtoNoNextHdr:
fp.Protocol = proto
fp.RemotePort = 0
fp.LocalPort = 0
fp.Fragment = false
return true
case ipProtoICMPv6:
if dataLen < offset+6 {
return false
}
fp.Protocol = proto
fp.LocalPort = 0
switch pkt[offset+1] {
case icmpv6TypeEchoRequest, icmpv6TypeEchoReply:
fp.RemotePort = binary.BigEndian.Uint16(pkt[offset+4 : offset+6])
default:
fp.RemotePort = 0
}
fp.Fragment = false
return true
case ipProtoTCP, ipProtoUDP:
if dataLen < offset+4 {
return false
}
fp.Protocol = proto
fp.RemotePort = binary.BigEndian.Uint16(pkt[offset : offset+2])
fp.LocalPort = binary.BigEndian.Uint16(pkt[offset+2 : offset+4])
fp.Fragment = false
return true
case ipProtoIPv6Fragment:
if dataLen < offset+8 {
return false
}
fragmentOffset := binary.BigEndian.Uint16(pkt[offset+2:offset+4]) &^ uint16(0x7)
if fragmentOffset != 0 {
fp.Protocol = pkt[offset]
fp.Fragment = true
fp.RemotePort = 0
fp.LocalPort = 0
return true
}
next = 8
case ipProtoAH:
if dataLen <= offset+1 {
return false
}
next = int(pkt[offset+1]+2) << 2
default:
if dataLen <= offset+1 {
return false
}
next = int(pkt[offset+1]+1) << 3
}
if next <= 0 {
next = 8
}
protoAt = offset
offset = offset + next
}
}
// runRxSplit drives the split path: faithful inbound parse for the firewall
// side, then m.Commit re-parses to coalesce. v6 controls which baseline
// parser we run.
func runRxSplit(b *testing.B, pkts [][]byte, batchSize int, v6 bool) {
b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, true, true)
var fp firewall.Packet
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)]
var ok bool
if v6 {
ok = parseV6InboundBaseline(pkt, &fp)
} else {
ok = parseV4InboundBaseline(pkt, &fp)
}
if !ok {
b.Fatal("baseline parse failed")
}
if err := m.Commit(pkt); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := m.Flush(); err != nil {
b.Fatal(err)
}
}
}
_ = m.Flush()
}
// runRxUnified drives the unified path: ParseInbound walks once, filling
// the conntrack key + coalescer hint in parsed; CommitInbound dispatches
// without re-parsing.
func runRxUnified(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, true, true)
var parsed RxParsed
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)]
if err := ParsePacket(pkt, true, &parsed); err != nil {
b.Fatal(err)
}
if err := m.CommitInbound(pkt, &parsed); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := m.Flush(); err != nil {
b.Fatal(err)
}
}
}
_ = m.Flush()
}
// buildUDPv4Bulk returns N UDP packets on a single 5-tuple suitable for the
// UDP coalescer's append path.
func buildUDPv4Bulk(n, payloadLen int) [][]byte {
pkts := make([][]byte, n)
pay := make([]byte, payloadLen)
for i := range n {
pkts[i] = buildUDPv4(1000, 53, pay)
}
return pkts
}
func buildTCPv6Bulk(n, payloadLen int) [][]byte {
pkts := make([][]byte, n)
pay := make([]byte, payloadLen)
seq := uint32(1000)
for i := range n {
pkts[i] = buildTCPv6(0, seq, tcpAck, pay)
seq += uint32(payloadLen)
}
return pkts
}
func buildICMPv4Bulk(n int) [][]byte {
pkts := make([][]byte, n)
for i := range pkts {
pkts[i] = buildICMPv4()
}
return pkts
}
// === TCPv4 ===
func BenchmarkRxSplitTCPv4(b *testing.B) {
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
runRxSplit(b, pkts, tcpCoalesceMaxSegs, false)
}
func BenchmarkRxUnifiedTCPv4(b *testing.B) {
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
runRxUnified(b, pkts, tcpCoalesceMaxSegs)
}
// === TCPv4 interleaved (4 flows) ===
func BenchmarkRxSplitTCPv4Interleaved4(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runRxSplit(b, pkts, len(pkts), false)
}
func BenchmarkRxUnifiedTCPv4Interleaved4(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runRxUnified(b, pkts, len(pkts))
}
// === UDPv4 ===
func BenchmarkRxSplitUDPv4(b *testing.B) {
pkts := buildUDPv4Bulk(udpCoalesceMaxSegs, 1200)
runRxSplit(b, pkts, udpCoalesceMaxSegs, false)
}
func BenchmarkRxUnifiedUDPv4(b *testing.B) {
pkts := buildUDPv4Bulk(udpCoalesceMaxSegs, 1200)
runRxUnified(b, pkts, udpCoalesceMaxSegs)
}
// === TCPv6 ===
func BenchmarkRxSplitTCPv6(b *testing.B) {
pkts := buildTCPv6Bulk(tcpCoalesceMaxSegs, 1200)
runRxSplit(b, pkts, tcpCoalesceMaxSegs, true)
}
func BenchmarkRxUnifiedTCPv6(b *testing.B) {
pkts := buildTCPv6Bulk(tcpCoalesceMaxSegs, 1200)
runRxUnified(b, pkts, tcpCoalesceMaxSegs)
}
// === ICMPv4 (passthrough) — measures the unified parser on the coalescer-
// rejected path, where both lenient and unified must still fill fp. ===
func BenchmarkRxSplitICMPv4(b *testing.B) {
pkts := buildICMPv4Bulk(64)
runRxSplit(b, pkts, 64, false)
}
func BenchmarkRxUnifiedICMPv4(b *testing.B) {
pkts := buildICMPv4Bulk(64)
runRxUnified(b, pkts, 64)
}
// === Firewall fast-path (conntrack-hit) — exercises the savings from the
// dense PacketKey: smaller hash key for the per-routine ConntrackCache,
// and skipping the AddrFrom4 calls that the old path needed to fill the
// netip.Addr-rich firewall.Packet up-front. ===
//
// The "split" baseline simulates the legacy path: parseV4InboundBaseline
// fills a netip.Addr-rich Packet, then we probe a localCache keyed on
// Packet. The "unified" path: ParseInbound fills only the dense PacketKey,
// and we probe a localCache keyed on PacketKey. Both paths follow with
// the coalescer Commit so the bench captures end-to-end RX-side cost.
// runRxSplitWithCache mirrors runRxSplit but runs the legacy-style
// firewall fast path (localCache keyed on firewall.Packet) on every
// packet so we can compare against the unified path.
func runRxSplitWithCache(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, true, true)
var fp firewall.Packet
// Pre-warm a per-packet cache keyed on the netip.Addr-rich Packet form.
cache := make(map[firewall.Packet]struct{}, len(pkts))
for _, pkt := range pkts {
var seedFp firewall.Packet
if !parseV4InboundBaseline(pkt, &seedFp) {
b.Fatal("seed parse failed")
}
cache[seedFp] = struct{}{}
}
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)]
if !parseV4InboundBaseline(pkt, &fp) {
b.Fatal("baseline parse failed")
}
if _, ok := cache[fp]; !ok {
b.Fatal("cache miss")
}
if err := m.Commit(pkt); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := m.Flush(); err != nil {
b.Fatal(err)
}
}
}
_ = m.Flush()
}
// runRxUnifiedWithCache: unified path with a PacketKey-keyed localCache.
// Each iteration: ParseInbound → conntrack-cache hit → CommitInbound.
func runRxUnifiedWithCache(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, true, true)
var parsed RxParsed
cache := make(firewall.ConntrackCache, len(pkts))
for _, pkt := range pkts {
var seed RxParsed
if err := ParsePacket(pkt, true, &seed); err != nil {
b.Fatal(err)
}
cache[seed.Key] = struct{}{}
}
b.ReportAllocs()
b.SetBytes(int64(len(pkts[0])))
b.ResetTimer()
for i := 0; i < b.N; i++ {
pkt := pkts[i%len(pkts)]
if err := ParsePacket(pkt, true, &parsed); err != nil {
b.Fatal(err)
}
if _, ok := cache[parsed.Key]; !ok {
b.Fatal("cache miss")
}
if err := m.CommitInbound(pkt, &parsed); err != nil {
b.Fatal(err)
}
if (i+1)%batchSize == 0 {
if err := m.Flush(); err != nil {
b.Fatal(err)
}
}
}
_ = m.Flush()
}
func BenchmarkRxSplitTCPv4WithCache(b *testing.B) {
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
runRxSplitWithCache(b, pkts, tcpCoalesceMaxSegs)
}
func BenchmarkRxUnifiedTCPv4WithCache(b *testing.B) {
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
runRxUnifiedWithCache(b, pkts, tcpCoalesceMaxSegs)
}
func BenchmarkRxSplitInterleaved4WithCache(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runRxSplitWithCache(b, pkts, len(pkts))
}
func BenchmarkRxUnifiedInterleaved4WithCache(b *testing.B) {
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
runRxUnifiedWithCache(b, pkts, len(pkts))
}
+174
View File
@@ -0,0 +1,174 @@
package batch
import (
"net/netip"
"testing"
"github.com/slackhq/nebula/firewall"
)
// TestParseInboundParity asserts that ParseInbound + Key.Hydrate produces
// the same firewall.Packet that the lenient baseline parsers (which
// mirror outside.go's parseV4/parseV6 with incoming=true) produce for
// every shape we care about. Catches drift between the unified
// parse-then-hydrate flow and the production newPacket behavior so
// swapping one for the other is observably safe.
func TestParseInboundParity(t *testing.T) {
cases := []struct {
name string
pkt []byte
v6 bool
}{
{"tcp_v4", buildTCPv4Ports(1234, 443, 1000, tcpAck, []byte("payload")), false},
{"tcp_v4_psh", buildTCPv4Ports(1234, 443, 2000, tcpAckPsh, make([]byte, 1200)), false},
{"udp_v4", buildUDPv4(40000, 53, []byte("dnsquery")), false},
{"icmp_v4", buildICMPv4(), false},
{"tcp_v6", buildTCPv6(0, 5000, tcpAck, make([]byte, 800)), true},
{"udp_v6", buildUDPv6(40001, 53, []byte("v6dns")), true},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
var fpUnified, fpBaseline firewall.Packet
var parsed RxParsed
if err := ParsePacket(tc.pkt, true, &parsed); err != nil {
t.Fatalf("ParsePacket: %v", err)
}
parsed.Key.Hydrate(&fpUnified)
var ok bool
if tc.v6 {
ok = parseV6InboundBaseline(tc.pkt, &fpBaseline)
} else {
ok = parseV4InboundBaseline(tc.pkt, &fpBaseline)
}
if !ok {
t.Fatalf("baseline parse failed")
}
if fpUnified != fpBaseline {
t.Errorf("firewall.Packet mismatch:\n unified: %+v\n baseline: %+v", fpUnified, fpBaseline)
}
})
}
}
// TestParseInboundFlowKey checks that the coalescer hint the unified parser
// produces matches what parseTCPBase/parseUDP would produce on the same
// packet — same flowKey, ipHdrLen, payLen, etc. The hint is only valid
// when Kind is RxKindTCP/RxKindUDP.
func TestParseInboundFlowKey(t *testing.T) {
t.Run("tcp_v4", func(t *testing.T) {
pkt := buildTCPv4Ports(1234, 443, 5000, tcpAck, make([]byte, 800))
var parsed RxParsed
if err := ParsePacket(pkt, true, &parsed); err != nil {
t.Fatal(err)
}
if parsed.Kind != RxKindTCP {
t.Fatalf("kind=%v want TCP", parsed.Kind)
}
ref, ok := parseTCPBase(pkt)
if !ok {
t.Fatal("parseTCPBase failed")
}
if parsed.tcp != ref {
t.Errorf("parsedTCP mismatch:\n unified: %+v\n ref: %+v", parsed.tcp, ref)
}
})
t.Run("udp_v4", func(t *testing.T) {
pkt := buildUDPv4(40000, 53, []byte("dnsquery"))
var parsed RxParsed
if err := ParsePacket(pkt, true, &parsed); err != nil {
t.Fatal(err)
}
if parsed.Kind != RxKindUDP {
t.Fatalf("kind=%v want UDP", parsed.Kind)
}
ref, ok := parseUDP(pkt)
if !ok {
t.Fatal("parseUDP failed")
}
if parsed.udp != ref {
t.Errorf("parsedUDP mismatch:\n unified: %+v\n ref: %+v", parsed.udp, ref)
}
})
t.Run("tcp_v6", func(t *testing.T) {
pkt := buildTCPv6(0, 9000, tcpAck, make([]byte, 800))
var parsed RxParsed
if err := ParsePacket(pkt, true, &parsed); err != nil {
t.Fatal(err)
}
if parsed.Kind != RxKindTCP {
t.Fatalf("kind=%v want TCP", parsed.Kind)
}
ref, ok := parseTCPBase(pkt)
if !ok {
t.Fatal("parseTCPBase failed")
}
if parsed.tcp != ref {
t.Errorf("parsedTCP mismatch:\n unified: %+v\n ref: %+v", parsed.tcp, ref)
}
})
}
// TestParseInboundICMPPassthrough confirms ICMP packets populate the
// conntrack key (including the ICMP identifier in RemotePort) but stay
// RxKindPassthrough so the batcher writes them verbatim. After Hydrate
// the firewall.Packet form should match what the legacy parseV4 produced.
func TestParseInboundICMPPassthrough(t *testing.T) {
pkt := buildICMPv4()
// Stamp a non-zero identifier into the ICMP header so we can check
// RemotePort gets it.
pkt[20] = 8 // type=echo
pkt[24] = 0xab
pkt[25] = 0xcd
var parsed RxParsed
if err := ParsePacket(pkt, true, &parsed); err != nil {
t.Fatal(err)
}
if parsed.Kind != RxKindPassthrough {
t.Errorf("kind=%v want Passthrough", parsed.Kind)
}
var fp firewall.Packet
parsed.Key.Hydrate(&fp)
if fp.Protocol != firewall.ProtoICMP {
t.Errorf("Protocol=%d want %d", fp.Protocol, firewall.ProtoICMP)
}
if fp.RemotePort != 0xabcd {
t.Errorf("RemotePort=0x%x want 0xabcd", fp.RemotePort)
}
if fp.LocalPort != 0 {
t.Errorf("LocalPort=%d want 0", fp.LocalPort)
}
wantRemote := netip.MustParseAddr("10.0.0.1")
wantLocal := netip.MustParseAddr("10.0.0.2")
if fp.RemoteAddr != wantRemote || fp.LocalAddr != wantLocal {
t.Errorf("addrs: remote=%v local=%v want %v/%v", fp.RemoteAddr, fp.LocalAddr, wantRemote, wantLocal)
}
}
// TestParseInboundV4Fragment confirms a fragmented v4 packet fills the
// conntrack key with Fragment=true and falls into Passthrough on the
// coalescer side.
func TestParseInboundV4Fragment(t *testing.T) {
// Build a TCP packet then twiddle the IP flags to make it look like a
// non-first fragment (offset != 0).
pkt := buildTCPv4Ports(1234, 443, 1000, tcpAck, []byte("payload"))
// Set a non-zero fragment offset (bytes 6-7, low 13 bits).
pkt[6] = 0x00
pkt[7] = 0x10 // offset = 16 (in 8-byte units)
var parsed RxParsed
if err := ParsePacket(pkt, true, &parsed); err != nil {
t.Fatal(err)
}
if !parsed.Key.Fragment {
t.Error("Fragment=false, want true")
}
if parsed.Kind != RxKindPassthrough {
t.Errorf("kind=%v want Passthrough", parsed.Kind)
}
}
+21 -20
View File
@@ -3,7 +3,6 @@ package batch
import ( import (
"errors" "errors"
"io" "io"
"log/slog"
) )
// MultiCoalescer fans plaintext packets out to lane-specific batchers based // MultiCoalescer fans plaintext packets out to lane-specific batchers based
@@ -27,38 +26,38 @@ type MultiCoalescer struct {
tcp *TCPCoalescer tcp *TCPCoalescer
udp *UDPCoalescer udp *UDPCoalescer
pt *Passthrough pt *Passthrough
// arena is owned by the Multi: lanes get only its Reserve (nil Resetter)
// and Flush resets it exactly once after every lane has drained.
arena *Arena
}
// DefaultMultiArenaCap is the recommended arena capacity for a Multi-lane // arena shared across all lanes so a single Reserve grows one backing
// batcher: 64 slots × 65535 bytes ≈ 4 MiB, enough to hold one recvmmsg // slice; lane Commit calls borrow into this same arena.
// burst worth of MTU-sized packets without the arena growing. backing []byte
const DefaultMultiArenaCap = initialSlots * 65535 }
// NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller // NewMultiCoalescer builds a multi-lane batcher. tcpEnabled lets the caller
// opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled // opt out of TCP coalescing (e.g. when the queue can't do TSO); udpEnabled
// likewise gates UDP coalescing (only enable when USO was negotiated). // likewise gates UDP coalescing (only enable when USO was negotiated).
// Either lane disabled redirects its traffic into the passthrough lane. // Either lane disabled redirects its traffic into the passthrough lane.
// arena is the single backing slab shared across every lane; the caller func NewMultiCoalescer(w io.Writer, tcpEnabled, udpEnabled bool) *MultiCoalescer {
// pre-sizes it via NewArena so the hot path never allocates.
func NewMultiCoalescer(w io.Writer, l *slog.Logger, arena *Arena, tcpEnabled, udpEnabled bool) *MultiCoalescer {
m := &MultiCoalescer{ m := &MultiCoalescer{
pt: NewPassthrough(w, arena.Reserve, nil), pt: NewPassthrough(w),
arena: arena, backing: make([]byte, 0, initialSlots*65535),
} }
if tcpEnabled { if tcpEnabled {
m.tcp = NewTCPCoalescer(w, l, arena.Reserve, nil) m.tcp = NewTCPCoalescer(w)
} }
if udpEnabled { if udpEnabled {
m.udp = NewUDPCoalescer(w, arena.Reserve, nil) m.udp = NewUDPCoalescer(w)
} }
return m return m
} }
func (m *MultiCoalescer) Reserve(sz int) []byte { func (m *MultiCoalescer) Reserve(sz int) []byte {
return m.arena.Reserve(sz) if len(m.backing)+sz > cap(m.backing) {
newCap := max(cap(m.backing)*2, sz)
m.backing = make([]byte, 0, newCap)
}
start := len(m.backing)
m.backing = m.backing[:start+sz]
return m.backing[start : start+sz : start+sz]
} }
// Commit dispatches pkt to the appropriate lane based on IP version + L4 // Commit dispatches pkt to the appropriate lane based on IP version + L4
@@ -110,8 +109,10 @@ func (m *MultiCoalescer) Commit(pkt []byte) error {
return m.pt.Commit(pkt) return m.pt.Commit(pkt)
} }
// Flush drains every lane in a fixed order, then resets the shared arena once. // Flush drains every lane in a fixed order: TCP, UDP, passthrough. Errors
// A lane error doesn't stop the remaining lanes; the joined errors are returned. // from a lane do not stop subsequent lanes from flushing, we keep
// draining and return the first observed error so a single bad packet
// doesn't strand the others.
func (m *MultiCoalescer) Flush() error { func (m *MultiCoalescer) Flush() error {
var errs []error var errs []error
if m.tcp != nil { if m.tcp != nil {
@@ -127,6 +128,6 @@ func (m *MultiCoalescer) Flush() error {
if err := m.pt.Flush(); err != nil { if err := m.pt.Flush(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
m.arena.Reset() m.backing = m.backing[:0]
return errors.Join(errs...) return errors.Join(errs...)
} }
+3 -5
View File
@@ -2,8 +2,6 @@ package batch
import ( import (
"testing" "testing"
"github.com/slackhq/nebula/test"
) )
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right // TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
@@ -11,7 +9,7 @@ import (
// else (ICMP here) falls through to plain Write. // else (ICMP here) falls through to plain Write.
func TestMultiCoalescerRoutesByProto(t *testing.T) { func TestMultiCoalescerRoutesByProto(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, true) m := NewMultiCoalescer(w, true, true)
tcpPay := make([]byte, 1200) tcpPay := make([]byte, 1200)
udpPay := make([]byte, 1200) udpPay := make([]byte, 1200)
@@ -53,7 +51,7 @@ func TestMultiCoalescerRoutesByProto(t *testing.T) {
// the kernel via the passthrough lane rather than being lost. // the kernel via the passthrough lane rather than being lost.
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) { func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), true, false) // TSO on, USO off m := NewMultiCoalescer(w, true, false) // TSO on, USO off
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil { if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -75,7 +73,7 @@ func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case. // TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) { func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
m := NewMultiCoalescer(w, test.NewLogger(), NewArena(0), false, true) // TSO off, USO on m := NewMultiCoalescer(w, false, true) // TSO off, USO on
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
+21 -22
View File
@@ -10,28 +10,29 @@ import (
type Passthrough struct { type Passthrough struct {
out io.Writer out io.Writer
slots [][]byte slots [][]byte
reserver Reserver backing []byte
resetter Resetter
cursor int cursor int
} }
const passthroughBaseNumSlots = 128 func NewPassthrough(w io.Writer) *Passthrough {
const baseNumSlots = 128
// DefaultPassthroughArenaCap is the recommended arena capacity for a
// standalone Passthrough batcher: 128 slots × udp.MTU ≈ 1.1 MiB.
const DefaultPassthroughArenaCap = passthroughBaseNumSlots * udp.MTU
func NewPassthrough(w io.Writer, reserver Reserver, resetter Resetter) *Passthrough {
return &Passthrough{ return &Passthrough{
out: w, out: w,
slots: make([][]byte, 0, passthroughBaseNumSlots), slots: make([][]byte, 0, baseNumSlots),
reserver: reserver, backing: make([]byte, 0, baseNumSlots*udp.MTU),
resetter: resetter,
} }
} }
func (p *Passthrough) Reserve(sz int) []byte { func (p *Passthrough) Reserve(sz int) []byte {
return p.reserver(sz) if len(p.backing)+sz > cap(p.backing) {
// Grow: allocate a fresh backing. Already-committed slices still
// reference the old array and remain valid until Flush drops them.
newCap := max(cap(p.backing)*2, sz)
p.backing = make([]byte, 0, newCap)
}
start := len(p.backing)
p.backing = p.backing[:start+sz]
return p.backing[start : start+sz : start+sz] //return zero length, sz-cap slice
} }
func (p *Passthrough) Commit(pkt []byte) error { func (p *Passthrough) Commit(pkt []byte) error {
@@ -39,17 +40,14 @@ func (p *Passthrough) Commit(pkt []byte) error {
return nil return nil
} }
// Flush drains every queued packet and calls the configured Resetter // CommitInbound ignores the hint — Passthrough never coalesces, so there's
func (p *Passthrough) Flush() error { // no IP/L4 re-parse to skip. Present so Passthrough satisfies the RxBatcher
firstErr := p.drain() // interface alongside MultiCoalescer.
if p.resetter != nil { func (p *Passthrough) CommitInbound(pkt []byte, _ *RxParsed) error {
p.resetter() return p.Commit(pkt)
}
return firstErr
} }
// drain writes out every queued packet and clears the slot list. func (p *Passthrough) Flush() error {
func (p *Passthrough) drain() error {
var firstErr error var firstErr error
for _, s := range p.slots { for _, s := range p.slots {
_, err := p.out.Write(s) _, err := p.out.Write(s)
@@ -59,5 +57,6 @@ func (p *Passthrough) drain() error {
} }
clear(p.slots) clear(p.slots)
p.slots = p.slots[:0] p.slots = p.slots[:0]
p.backing = p.backing[:0]
return firstErr return firstErr
} }
+42 -39
View File
@@ -2,7 +2,6 @@ package batch
import ( import (
"bytes" "bytes"
"context"
"encoding/binary" "encoding/binary"
"io" "io"
"log/slog" "log/slog"
@@ -12,7 +11,10 @@ import (
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
) )
// ipProtoTCP is the IANA protocol number for TCP. Defined here to help Windows out. // ipProtoTCP is the IANA protocol number for TCP. Hardcoded instead of
// reaching for golang.org/x/sys/unix — that package doesn't define the
// constant on Windows, which would break cross-compiles even though this
// file runs unchanged on every platform.
const ipProtoTCP = 6 const ipProtoTCP = 6
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's // tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
@@ -20,7 +22,8 @@ const ipProtoTCP = 6
const tcpCoalesceBufSize = 65535 const tcpCoalesceBufSize = 65535
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single // tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
// superpacket. Keeping this well below the kernel's TSO ceiling bounds latency. // superpacket. Keeping this well below the kernel's TSO ceiling bounds
// latency.
const tcpCoalesceMaxSegs = 64 const tcpCoalesceMaxSegs = 64
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header // tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
@@ -79,20 +82,17 @@ type TCPCoalescer struct {
// at is removed/sealed. // at is removed/sealed.
lastSlot *coalesceSlot lastSlot *coalesceSlot
pool []*coalesceSlot // free list for reuse pool []*coalesceSlot // free list for reuse
reserver Reserver
resetter Resetter backing []byte
l *slog.Logger
} }
func NewTCPCoalescer(w io.Writer, l *slog.Logger, reserver Reserver, resetter Resetter) *TCPCoalescer { func NewTCPCoalescer(w io.Writer) *TCPCoalescer {
c := &TCPCoalescer{ c := &TCPCoalescer{
plainW: w, plainW: w,
slots: make([]*coalesceSlot, 0, initialSlots), slots: make([]*coalesceSlot, 0, initialSlots),
openSlots: make(map[flowKey]*coalesceSlot, initialSlots), openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
pool: make([]*coalesceSlot, 0, initialSlots), pool: make([]*coalesceSlot, 0, initialSlots),
reserver: reserver, backing: make([]byte, 0, initialSlots*65535),
resetter: resetter,
l: l,
} }
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok { if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
c.gsoW = gw c.gsoW = gw
@@ -114,8 +114,8 @@ type parsedTCP struct {
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet, // parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
// regardless of whether it's admissible for coalescing. Returns ok=false // regardless of whether it's admissible for coalescing. Returns ok=false
// for non-TCP or malformed input. // for non-TCP or malformed input. Accepts IPv4 (no options, no fragmentation)
// Accepts IPv4 (no options or fragmentation) and IPv6 (no extension headers). // and IPv6 (no extension headers).
func parseTCPBase(pkt []byte) (parsedTCP, bool) { func parseTCPBase(pkt []byte) (parsedTCP, bool) {
var p parsedTCP var p parsedTCP
ip, ok := parseIPPrologue(pkt, ipProtoTCP) ip, ok := parseIPPrologue(pkt, ipProtoTCP)
@@ -171,10 +171,12 @@ func (p parsedTCP) coalesceable() bool {
} }
func (c *TCPCoalescer) Reserve(sz int) []byte { func (c *TCPCoalescer) Reserve(sz int) []byte {
return c.reserver(sz) return reserveFromBacking(&c.backing, sz)
} }
// Commit borrows pkt. The caller must keep pkt valid until the next Flush. // Commit borrows pkt. The caller must keep pkt valid until the next Flush,
// whether or not the packet was coalesced — passthrough (non-admissible)
// packets are queued and written at Flush time, not synchronously.
func (c *TCPCoalescer) Commit(pkt []byte) error { func (c *TCPCoalescer) Commit(pkt []byte) error {
if c.gsoW == nil { if c.gsoW == nil {
c.addPassthrough(pkt) c.addPassthrough(pkt)
@@ -240,18 +242,16 @@ func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
return nil return nil
} }
// Flush emits every queued event in (per-flow) seq order. // Flush emits every queued event in (per-flow) seq order. Coalesced slots
// go out via WriteGSO; passthrough slots go out via plainW.Write.
// reorderForFlush first sorts each flow's slots into TCP-seq order within
// passthrough-bounded segments and merges contiguous adjacent slots, so
// any wire-side reorder that crossed an rxOrder batch boundary doesn't
// get amplified into kernel-visible reorder by the slot machinery.
// Returns the first error observed; keeps draining so one bad packet
// doesn't hold up the rest. After Flush returns, borrowed payload slices
// may be recycled.
func (c *TCPCoalescer) Flush() error { func (c *TCPCoalescer) Flush() error {
first := c.drain()
if c.resetter != nil {
c.resetter()
}
return first
}
// drain emits every queued slot (reordering/merging coalesced runs first)
// and clears the slot state.
func (c *TCPCoalescer) drain() error {
c.reorderForFlush() c.reorderForFlush()
var first error var first error
for _, s := range c.slots { for _, s := range c.slots {
@@ -271,6 +271,7 @@ func (c *TCPCoalescer) drain() error {
clear(c.openSlots) clear(c.openSlots)
c.lastSlot = nil c.lastSlot = nil
c.backing = c.backing[:0]
return first return first
} }
@@ -314,7 +315,8 @@ func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
} }
// canAppend reports whether info's packet extends the slot's seed: same // canAppend reports whether info's packet extends the slot's seed: same
// header shape and stable contents, adjacent seq, not oversized, chain not closed. // header shape and stable contents, adjacent seq, not oversized, chain not
// closed.
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool { func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
if s.psh { if s.psh {
return false return false
@@ -356,6 +358,9 @@ func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP
// last segment. Without this the sender's push signal is dropped. // last segment. Without this the sender's push signal is dropped.
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
} }
// Merge IP-level CE marks into the seed: headersMatch ignores ECN, so
// this is the one place the signal is preserved.
mergeECNIntoSeed(s.hdrBuf[:s.ipHdrLen], pkt[:s.ipHdrLen], s.isV6)
if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 { if info.payLen < s.gsoSize || info.flags&tcpFlagPsh != 0 {
s.psh = true s.psh = true
} }
@@ -382,7 +387,8 @@ func (c *TCPCoalescer) release(s *coalesceSlot) {
c.pool = append(c.pool, s) c.pool = append(c.pool, s)
} }
// flushSlot patches the header and calls WriteGSO. Does not remove the slot from c.slots. // flushSlot patches the header and calls WriteGSO. Does not remove the
// slot from c.slots.
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error { func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
total := s.hdrLen + s.totalPay total := s.hdrLen + s.totalPay
l4Len := total - s.ipHdrLen l4Len := total - s.ipHdrLen
@@ -411,7 +417,8 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
// headersMatch compares two IP+TCP header prefixes for byte-for-byte // headersMatch compares two IP+TCP header prefixes for byte-for-byte
// equality on every field that must be identical across coalesced // equality on every field that must be identical across coalesced
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out. // segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out, as is the
// 2-bit IP-level ECN field — appendPayload merges CE into the seed.
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool { func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if len(a) != len(b) { if len(a) != len(b) {
return false return false
@@ -481,11 +488,10 @@ func (c *TCPCoalescer) reorderForFlush() {
// the operator can quantify how often it happens; the data // the operator can quantify how often it happens; the data
// itself still emits in seq order, kernel TCP handles the // itself still emits in seq order, kernel TCP handles the
// gap via its OOO queue. // gap via its OOO queue.
if c.l.Enabled(context.Background(), slog.LevelDebug) {
if prev.nextSeq != slotSeedSeq(s) { if prev.nextSeq != slotSeedSeq(s) {
logged = true logged = true
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq) gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
c.l.Debug("tcp coalesce: cross-slot seq gap", slog.Default().Warn("tcp coalesce: cross-slot seq gap",
"src", flowKeyAddr(s.fk, false), "src", flowKeyAddr(s.fk, false),
"dst", flowKeyAddr(s.fk, true), "dst", flowKeyAddr(s.fk, true),
"sport", s.fk.sport, "sport", s.fk.sport,
@@ -498,8 +504,6 @@ func (c *TCPCoalescer) reorderForFlush() {
"prev_total_pay", prev.totalPay, "prev_total_pay", prev.totalPay,
) )
} }
}
if canMergeSlots(prev, s) { if canMergeSlots(prev, s) {
mergeSlots(prev, s) mergeSlots(prev, s)
c.release(s) c.release(s)
@@ -510,7 +514,7 @@ func (c *TCPCoalescer) reorderForFlush() {
out = append(out, s) out = append(out, s)
} }
if logged { if logged {
c.l.Warn("==== end of batch ====") slog.Default().Warn("==== end of batch ====")
} }
c.slots = out c.slots = out
} }
@@ -618,9 +622,6 @@ func flowKeyCompare(a, b flowKey) int {
// ECE state must agree across both slots: PSH is a semantic delimiter // ECE state must agree across both slots: PSH is a semantic delimiter
// (preserving the sender's push boundary) and ECE state must be uniform // (preserving the sender's push boundary) and ECE state must be uniform
// across a window (the same rule canAppend enforces for in-flow appends). // across a window (the same rule canAppend enforces for in-flow appends).
// The IP-level ECN codepoint must also match: this check calls headersMatch
// → ipHeadersMatch, which compares the full DSCP/ECN byte, so two slots with
// differing ECN marks stay separate superpackets, each keeping its own mark.
// //
// Note: a slot sealed by reorder (canAppend returned false on seq // Note: a slot sealed by reorder (canAppend returned false on seq
// mismatch) keeps psh=false, so this restriction does not block the // mismatch) keeps psh=false, so this restriction does not block the
@@ -659,9 +660,10 @@ func canMergeSlots(prev, s *coalesceSlot) bool {
} }
// mergeSlots folds src into dst in place: payIovs concatenated, counters // mergeSlots folds src into dst in place: payIovs concatenated, counters
// and totals updated, PSH OR'd into the seed header so the push signal is // and totals updated, PSH and IP-level CE bits OR'd into the seed header
// not lost. The seed header's seq, gsoSize, and fk are unchanged. Caller // so neither the push signal nor a CE mark is lost. The seed header's
// is responsible for releasing src (it's no longer in c.slots after this call). // seq, gsoSize, and fk are unchanged. Caller is responsible for releasing
// src (it's no longer in c.slots after this call).
func mergeSlots(dst, src *coalesceSlot) { func mergeSlots(dst, src *coalesceSlot) {
dst.payIovs = append(dst.payIovs, src.payIovs...) dst.payIovs = append(dst.payIovs, src.payIovs...)
dst.numSeg += src.numSeg dst.numSeg += src.numSeg
@@ -671,6 +673,7 @@ func mergeSlots(dst, src *coalesceSlot) {
dst.psh = true dst.psh = true
dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh dst.hdrBuf[dst.ipHdrLen+13] |= tcpFlagPsh
} }
mergeECNIntoSeed(dst.hdrBuf[:dst.ipHdrLen], src.hdrBuf[:src.ipHdrLen], dst.isV6)
} }
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must // ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
+2 -4
View File
@@ -6,7 +6,6 @@ import (
"testing" "testing"
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test"
) )
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards // nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
@@ -71,8 +70,7 @@ func buildICMPv4() []byte {
// between batches, and reports per-packet cost. // between batches, and reports per-packet cost.
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) { func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper() b.Helper()
arena := NewArena(0) c := NewTCPCoalescer(nopTunWriter{})
c := NewTCPCoalescer(nopTunWriter{}, test.NewLogger(), arena.Reserve, arena.Reset)
b.ReportAllocs() b.ReportAllocs()
b.SetBytes(int64(len(pkts[0]))) b.SetBytes(int64(len(pkts[0])))
b.ResetTimer() b.ResetTimer()
@@ -141,7 +139,7 @@ func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
// is the bench that shows the savings of skipping the lane's re-parse. // is the bench that shows the savings of skipping the lane's re-parse.
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) { func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
b.Helper() b.Helper()
m := NewMultiCoalescer(nopTunWriter{}, test.NewLogger(), NewArena(0), true, true) m := NewMultiCoalescer(nopTunWriter{}, true, true)
b.ReportAllocs() b.ReportAllocs()
b.SetBytes(int64(len(pkts[0]))) b.SetBytes(int64(len(pkts[0])))
b.ResetTimer() b.ResetTimer()
+50 -145
View File
@@ -5,7 +5,6 @@ import (
"testing" "testing"
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
"github.com/slackhq/nebula/test"
) )
// fakeTunWriter records plain Writes and WriteGSO calls without touching a // fakeTunWriter records plain Writes and WriteGSO calls without touching a
@@ -128,8 +127,7 @@ const (
func TestCoalescerPassthroughWhenGSOUnavailable(t *testing.T) { func TestCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: false} w := &fakeTunWriter{gsoEnabled: false}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pkt := buildTCPv4(1000, tcpAck, []byte("hello")) pkt := buildTCPv4(1000, tcpAck, []byte("hello"))
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -148,8 +146,7 @@ func TestCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
func TestCoalescerNonTCPPassthrough(t *testing.T) { func TestCoalescerNonTCPPassthrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pkt := make([]byte, 28) pkt := make([]byte, 28)
pkt[0] = 0x45 pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], 28) binary.BigEndian.PutUint16(pkt[2:4], 28)
@@ -169,8 +166,7 @@ func TestCoalescerNonTCPPassthrough(t *testing.T) {
func TestCoalescerSeedThenFlushAlone(t *testing.T) { func TestCoalescerSeedThenFlushAlone(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pkt := buildTCPv4(1000, tcpAck, make([]byte, 1000)) pkt := buildTCPv4(1000, tcpAck, make([]byte, 1000))
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -197,8 +193,7 @@ func TestCoalescerSeedThenFlushAlone(t *testing.T) {
func TestCoalescerCoalescesAdjacentACKs(t *testing.T) { func TestCoalescerCoalescesAdjacentACKs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -238,8 +233,7 @@ func TestCoalescerCoalescesAdjacentACKs(t *testing.T) {
func TestCoalescerRejectsSeqGap(t *testing.T) { func TestCoalescerRejectsSeqGap(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -258,8 +252,7 @@ func TestCoalescerRejectsSeqGap(t *testing.T) {
func TestCoalescerRejectsFlagMismatch(t *testing.T) { func TestCoalescerRejectsFlagMismatch(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -280,8 +273,7 @@ func TestCoalescerRejectsFlagMismatch(t *testing.T) {
func TestCoalescerRejectsFIN(t *testing.T) { func TestCoalescerRejectsFIN(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
fin := buildTCPv4(1000, tcpAck|tcpFin, []byte("x")) fin := buildTCPv4(1000, tcpAck|tcpFin, []byte("x"))
if err := c.Commit(fin); err != nil { if err := c.Commit(fin); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -297,8 +289,7 @@ func TestCoalescerRejectsFIN(t *testing.T) {
func TestCoalescerShortLastSegmentClosesChain(t *testing.T) { func TestCoalescerShortLastSegmentClosesChain(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
full := make([]byte, 1200) full := make([]byte, 1200)
half := make([]byte, 500) half := make([]byte, 500)
if err := c.Commit(buildTCPv4(1000, tcpAck, full)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, full)); err != nil {
@@ -333,8 +324,7 @@ func TestCoalescerShortLastSegmentClosesChain(t *testing.T) {
func TestCoalescerPSHFinalizesChain(t *testing.T) { func TestCoalescerPSHFinalizesChain(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -364,8 +354,7 @@ func TestCoalescerPSHFinalizesChain(t *testing.T) {
// coalescer drops it the sender's push signal never reaches the receiver. // coalescer drops it the sender's push signal never reaches the receiver.
func TestCoalescerPropagatesPSHFromAppended(t *testing.T) { func TestCoalescerPropagatesPSHFromAppended(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Seed has no PSH; second segment carries PSH and seals the chain. // Seed has no PSH; second segment carries PSH and seals the chain.
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
@@ -393,8 +382,7 @@ func TestCoalescerPropagatesPSHFromAppended(t *testing.T) {
func TestCoalescerRejectsDifferentFlow(t *testing.T) { func TestCoalescerRejectsDifferentFlow(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
p1 := buildTCPv4(1000, tcpAck, pay) p1 := buildTCPv4(1000, tcpAck, pay)
p2 := buildTCPv4(2200, tcpAck, pay) p2 := buildTCPv4(2200, tcpAck, pay)
@@ -416,8 +404,7 @@ func TestCoalescerRejectsDifferentFlow(t *testing.T) {
func TestCoalescerRejectsIPOptions(t *testing.T) { func TestCoalescerRejectsIPOptions(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 500) pay := make([]byte, 500)
pkt := buildTCPv4(1000, tcpAck, pay) pkt := buildTCPv4(1000, tcpAck, pay)
// Bump IHL to 6 to simulate 4 bytes of IP options. Don't actually add // Bump IHL to 6 to simulate 4 bytes of IP options. Don't actually add
@@ -437,8 +424,7 @@ func TestCoalescerRejectsIPOptions(t *testing.T) {
func TestCoalescerCapBySegments(t *testing.T) { func TestCoalescerCapBySegments(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 512) pay := make([]byte, 512)
seq := uint32(1000) seq := uint32(1000)
for i := 0; i < tcpCoalesceMaxSegs+5; i++ { for i := 0; i < tcpCoalesceMaxSegs+5; i++ {
@@ -462,8 +448,7 @@ func TestCoalescerCapBySegments(t *testing.T) {
// flows coalesce independently in a single Flush. // flows coalesce independently in a single Flush.
func TestCoalescerMultipleFlowsInSameBatch(t *testing.T) { func TestCoalescerMultipleFlowsInSameBatch(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Flow A: sport 1000. Flow B: sport 3000. // Flow A: sport 1000. Flow B: sport 3000.
@@ -520,8 +505,7 @@ func TestCoalescerMultipleFlowsInSameBatch(t *testing.T) {
// writing passthrough packets synchronously. // writing passthrough packets synchronously.
func TestCoalescerPreservesArrivalOrder(t *testing.T) { func TestCoalescerPreservesArrivalOrder(t *testing.T) {
w := &orderedFakeWriter{gsoEnabled: true} w := &orderedFakeWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
// Sequence: coalesceable TCP, ICMP (passthrough), coalesceable TCP on // Sequence: coalesceable TCP, ICMP (passthrough), coalesceable TCP on
// a different flow. Expected emit order: gso(X), plain(ICMP), gso(Y). // a different flow. Expected emit order: gso(X), plain(ICMP), gso(Y).
pay := make([]byte, 1200) pay := make([]byte, 1200)
@@ -589,8 +573,7 @@ func stringSliceEq(a, b []string) bool {
// packet (SYN) mid-flow only flushes its own flow, not others. // packet (SYN) mid-flow only flushes its own flow, not others.
func TestCoalescerInterleavedFlowsPreserveOrdering(t *testing.T) { func TestCoalescerInterleavedFlowsPreserveOrdering(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Flow A two segments. // Flow A two segments.
@@ -695,8 +678,7 @@ func buildTCPv6(tcLow byte, seq uint32, flags byte, payload []byte) []byte {
// retains ECE on the wire. // retains ECE on the wire.
func TestCoalescerCoalescesEceFlow(t *testing.T) { func TestCoalescerCoalescesEceFlow(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
flags := byte(tcpAck | tcpEce) flags := byte(tcpAck | tcpEce)
if err := c.Commit(buildTCPv4(1000, flags, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, flags, pay)); err != nil {
@@ -725,8 +707,7 @@ func TestCoalescerCoalescesEceFlow(t *testing.T) {
// in-flow segment seeds a new slot rather than extending the prior burst. // in-flow segment seeds a new slot rather than extending the prior burst.
func TestCoalescerCwrSealsFlow(t *testing.T) { func TestCoalescerCwrSealsFlow(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -759,8 +740,7 @@ func TestCoalescerCwrSealsFlow(t *testing.T) {
// a CE-echoing window or none. // a CE-echoing window or none.
func TestCoalescerEceMismatchReseeds(t *testing.T) { func TestCoalescerEceMismatchReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4(1000, tcpAck|tcpEce, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck|tcpEce, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -781,94 +761,42 @@ func TestCoalescerEceMismatchReseeds(t *testing.T) {
} }
} }
// TestCoalescerDifferingECNReseeds confirms that segments with differing IP // TestCoalescerMergesCEMark confirms that an ECT(0) burst with a single
// ECN codepoints do NOT coalesce: headersMatch compares the full ToS byte, // CE-marked packet still coalesces, and the merged superpacket carries CE.
// matching kernel GRO. Two ECT(0) segments merge; a CE stamp mid-run seals func TestCoalescerMergesCEMark(t *testing.T) {
// the ECT(0) chain and starts a fresh superpacket that keeps CE; a trailing
// ECT(0) starts yet another. Each superpacket keeps its own codepoint —
// ORing the marks (the old buggy behavior) would have fabricated a false CE
// across the whole burst.
func TestCoalescerDifferingECNReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4WithToS(ecnECT0, 1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4WithToS(ecnECT0, 1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := c.Commit(buildTCPv4WithToS(ecnECT0, 2200, tcpAck, pay)); err != nil {
t.Fatal(err)
}
// Router along the path stamped CE on this one. // Router along the path stamped CE on this one.
if err := c.Commit(buildTCPv4WithToS(ecnCE, 3400, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4WithToS(ecnCE, 2200, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := c.Commit(buildTCPv4WithToS(ecnECT0, 4600, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4WithToS(ecnECT0, 3400, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 3 { if len(w.gsoWrites) != 1 {
t.Fatalf("want 3 superpackets (ECN split), got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) t.Fatalf("want 1 merged gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
} }
// gso[0]: the two ECT(0) segments merged; gso[1]: CE alone; gso[2]: g := w.gsoWrites[0]
// trailing ECT(0) alone. Emitted in seq order. if len(g.pays) != 3 {
type want struct { t.Errorf("pay count=%d want 3", len(g.pays))
pays int
ecn byte
}
wants := []want{{2, ecnECT0}, {1, ecnCE}, {1, ecnECT0}}
for i, wnt := range wants {
g := w.gsoWrites[i]
if len(g.pays) != wnt.pays {
t.Errorf("gso %d pay count=%d want %d", i, len(g.pays), wnt.pays)
}
if got := g.hdr[1] & 0x03; got != wnt.ecn {
t.Errorf("gso %d ECN=0x%02x want 0x%02x", i, got, wnt.ecn)
} }
if got := g.hdr[1] & 0x03; got != ecnCE {
t.Errorf("seed ECN=0x%02x want CE 0x%02x", got, ecnCE)
} }
} }
// TestCoalescerECT0ThenECT1NoCE is the core regression for the ECN merge // TestCoalescerDscpMismatchReseeds confirms that the new ECN-mask in
// bug: ORing ECT(0)=0b10 with ECT(1)=0b01 fabricates CE=0b11. The two // headersMatch did not also relax DSCP — different DSCP must still split.
// segments must land in separate superpackets, each preserving its own
// codepoint, and neither may end up CE-marked.
func TestCoalescerECT0ThenECT1NoCE(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200)
if err := c.Commit(buildTCPv4WithToS(ecnECT0, 1000, tcpAck, pay)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildTCPv4WithToS(ecnECT1, 2200, tcpAck, pay)); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.gsoWrites) != 2 {
t.Fatalf("want 2 separate superpackets (ECT0 vs ECT1), got %d", len(w.gsoWrites))
}
wantECN := []byte{ecnECT0, ecnECT1}
for i, g := range w.gsoWrites {
if got := g.hdr[1] & 0x03; got != wantECN[i] {
t.Errorf("gso %d ECN=0x%02x want 0x%02x", i, got, wantECN[i])
}
if got := g.hdr[1] & 0x03; got == ecnCE {
t.Errorf("gso %d fabricated CE from ECT merge", i)
}
}
}
// TestCoalescerDscpMismatchReseeds confirms that a DSCP difference (same
// ECN) still splits — headersMatch compares the full ToS byte, so the upper
// six DSCP bits must match too.
func TestCoalescerDscpMismatchReseeds(t *testing.T) { func TestCoalescerDscpMismatchReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Same ECN (Not-ECT), different DSCP (0x10 vs 0x20 in upper 6 bits). // Same ECN (Not-ECT), different DSCP (0x10 vs 0x20 in upper 6 bits).
tosA := byte(0x10<<2) | ecnNotECT tosA := byte(0x10<<2) | ecnNotECT
@@ -891,8 +819,7 @@ func TestCoalescerDscpMismatchReseeds(t *testing.T) {
// TestCoalescerCoalescesEceFlow. // TestCoalescerCoalescesEceFlow.
func TestCoalescerIPv6CoalescesEceFlow(t *testing.T) { func TestCoalescerIPv6CoalescesEceFlow(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
flags := byte(tcpAck | tcpEce) flags := byte(tcpAck | tcpEce)
if err := c.Commit(buildTCPv6(0, 1000, flags, pay)); err != nil { if err := c.Commit(buildTCPv6(0, 1000, flags, pay)); err != nil {
@@ -923,8 +850,7 @@ func TestCoalescerIPv6CoalescesEceFlow(t *testing.T) {
// seen had the wire never reordered. // seen had the wire never reordered.
func TestCoalescerSortsReorderedSeedsAndMerges(t *testing.T) { func TestCoalescerSortsReorderedSeedsAndMerges(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Arrival order: seq 1000, 3400, 2200. The 3400 seeds a separate slot // Arrival order: seq 1000, 3400, 2200. The 3400 seeds a separate slot
// because 3400 != nextSeq=2200, then 2200 fails to extend the 3400 slot // because 3400 != nextSeq=2200, then 2200 fails to extend the 3400 slot
@@ -960,8 +886,7 @@ func TestCoalescerSortsReorderedSeedsAndMerges(t *testing.T) {
// without any cross-flow contamination. // without any cross-flow contamination.
func TestCoalescerSortAcrossFlowsMergesEachIndependently(t *testing.T) { func TestCoalescerSortAcrossFlowsMergesEachIndependently(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Flow A (sport 1000) seq 100, 1300; flow B (sport 3000) seq 500, 1700. // Flow A (sport 1000) seq 100, 1300; flow B (sport 3000) seq 500, 1700.
// Arrival: A.1300, B.1700, A.100, B.500 — every flow reordered. // Arrival: A.1300, B.1700, A.100, B.500 — every flow reordered.
@@ -1012,8 +937,7 @@ func TestCoalescerSortAcrossFlowsMergesEachIndependently(t *testing.T) {
// boundary by an arbitrary number of segments. // boundary by an arbitrary number of segments.
func TestCoalescerSortKeepsPSHBoundary(t *testing.T) { func TestCoalescerSortKeepsPSHBoundary(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// Seq 1000 (no PSH) + 2200 (PSH) → seal one slot with PSH set. // Seq 1000 (no PSH) + 2200 (PSH) → seal one slot with PSH set.
// Seq 3400 (no PSH) is contiguous to 3400 from seq 2200+1200; without // Seq 3400 (no PSH) is contiguous to 3400 from seq 2200+1200; without
@@ -1041,8 +965,7 @@ func TestCoalescerSortKeepsPSHBoundary(t *testing.T) {
// is sorted/merged independently. // is sorted/merged independently.
func TestCoalescerSortKeepsPassthroughBarrier(t *testing.T) { func TestCoalescerSortKeepsPassthroughBarrier(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// First two segments seed S1 (then a 3400 reorder seeds S2). // First two segments seed S1 (then a 3400 reorder seeds S2).
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
@@ -1071,48 +994,30 @@ func TestCoalescerSortKeepsPassthroughBarrier(t *testing.T) {
} }
} }
// TestCoalescerIPv6DifferingECNReseeds is the IPv6 analogue of // TestCoalescerIPv6MergesCEMark is the IPv6 analogue of
// TestCoalescerDifferingECNReseeds. ECN bits live in TC[1:0] = byte 1 mask // TestCoalescerMergesCEMark. ECN bits live in TC[1:0] = byte 1 mask 0x30.
// 0x30, so ipHeadersMatch (comparing byte 1 fully) still splits them. func TestCoalescerIPv6MergesCEMark(t *testing.T) {
func TestCoalescerIPv6DifferingECNReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewTCPCoalescer(w)
c := NewTCPCoalescer(w, test.NewLogger(), arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
// tcLow is the low 4 bits of TC; ECN occupies the bottom 2 of those. // tcLow is the low 4 bits of TC; ECN occupies the bottom 2 of those.
if err := c.Commit(buildTCPv6(ecnECT0, 1000, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv6(ecnECT0, 1000, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := c.Commit(buildTCPv6(ecnECT0, 2200, tcpAck, pay)); err != nil { if err := c.Commit(buildTCPv6(ecnCE, 2200, tcpAck, pay)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildTCPv6(ecnCE, 3400, tcpAck, pay)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildTCPv6(ecnECT0, 4600, tcpAck, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 3 { if len(w.gsoWrites) != 1 {
t.Fatalf("want 3 superpackets (ECN split), got %d", len(w.gsoWrites)) t.Fatalf("want 1 merged gso write, got %d", len(w.gsoWrites))
} }
g := w.gsoWrites[0]
// Byte 1 high nibble holds TC[3:0]; ECN is the low 2 bits of that nibble, // Byte 1 high nibble holds TC[3:0]; ECN is the low 2 bits of that nibble,
// which appears in byte 1 mask 0x30 (>>4 to read the codepoint value). // which appears in byte 1 mask 0x30 (>>4 to read the codepoint value).
type want struct { if got := (g.hdr[1] >> 4) & 0x03; got != ecnCE {
pays int t.Errorf("seed v6 ECN=0x%02x want CE 0x%02x", got, ecnCE)
ecn byte
}
wants := []want{{2, ecnECT0}, {1, ecnCE}, {1, ecnECT0}}
for i, wnt := range wants {
g := w.gsoWrites[i]
if len(g.pays) != wnt.pays {
t.Errorf("gso %d pay count=%d want %d", i, len(g.pays), wnt.pays)
}
if got := (g.hdr[1] >> 4) & 0x03; got != wnt.ecn {
t.Errorf("gso %d v6 ECN=0x%02x want 0x%02x", i, got, wnt.ecn)
}
} }
} }
+16 -12
View File
@@ -11,34 +11,38 @@ type batchWriter interface {
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch. // SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
// One SendBatch is owned by each listenIn goroutine; no locking is needed. // One SendBatch is owned by each listenIn goroutine; no locking is needed.
// Slots are backed by an Arena (see its docs) // The backing arena grows on demand: when there isn't room for the next slot
// we allocate a fresh backing array. Already-committed slices keep referencing
// the old array and remain valid until Flush drops them.
type SendBatch struct { type SendBatch struct {
out batchWriter out batchWriter
bufs [][]byte bufs [][]byte
dsts []netip.AddrPort dsts []netip.AddrPort
ecns []byte ecns []byte
arena *Arena backing []byte
} }
// NewSendBatch makes a SendBatch with batchCap slots and an arenaSize byte buffer for slices to back those slots func NewSendBatch(out batchWriter, batchCap, slotCap int) *SendBatch {
func NewSendBatch(out batchWriter, batchCap, arenaSize int) *SendBatch {
return &SendBatch{ return &SendBatch{
out: out, out: out,
bufs: make([][]byte, 0, batchCap), bufs: make([][]byte, 0, batchCap),
dsts: make([]netip.AddrPort, 0, batchCap), dsts: make([]netip.AddrPort, 0, batchCap),
ecns: make([]byte, 0, batchCap), ecns: make([]byte, 0, batchCap),
arena: NewArena(arenaSize), backing: make([]byte, 0, batchCap*slotCap),
} }
} }
func (b *SendBatch) Reserve(sz int) []byte { func (b *SendBatch) Reserve(sz int) []byte {
return b.arena.Reserve(sz) if len(b.backing)+sz > cap(b.backing) {
// Grow: allocate a fresh backing. Already-committed slices still
// reference the old array and remain valid until Flush drops them.
newCap := max(cap(b.backing)*2, sz)
b.backing = make([]byte, 0, newCap)
}
start := len(b.backing)
b.backing = b.backing[:start+sz]
return b.backing[start : start+sz : start+sz]
} }
// Len reports how many packets are queued for the next Flush. Callers use
// it to flush incrementally once a full sendmmsg batch has accumulated,
// bounding how long the first packet of a large read batch waits.
func (b *SendBatch) Len() int { return len(b.bufs) }
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) { func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
b.bufs = append(b.bufs, pkt) b.bufs = append(b.bufs, pkt)
@@ -55,6 +59,6 @@ func (b *SendBatch) Flush() error {
b.bufs = b.bufs[:0] b.bufs = b.bufs[:0]
b.dsts = b.dsts[:0] b.dsts = b.dsts[:0]
b.ecns = b.ecns[:0] b.ecns = b.ecns[:0]
b.arena.Reset() b.backing = b.backing[:0]
return err return err
} }
+14 -33
View File
@@ -22,7 +22,10 @@ const udpCoalesceMaxSegs = 64
// into. IPv6 (40) + UDP (8) = 48; round up for safety. // into. IPv6 (40) + UDP (8) = 48; round up for safety.
const udpCoalesceHdrCap = 64 const udpCoalesceHdrCap = 64
// udpSlot is one entry in the UDPCoalescer's ordered event queue. // udpSlot is one entry in the UDPCoalescer's ordered event queue. Same
// passthrough-vs-coalesced shape as the TCP coalescer's slot, but no
// seq/PSH/CWR bookkeeping — UDP segments only need 5-tuple + length
// matching to coalesce.
type udpSlot struct { type udpSlot struct {
passthrough bool passthrough bool
rawPkt []byte // borrowed when passthrough rawPkt []byte // borrowed when passthrough
@@ -62,8 +65,8 @@ type UDPCoalescer struct {
slots []*udpSlot slots []*udpSlot
openSlots map[flowKey]*udpSlot openSlots map[flowKey]*udpSlot
pool []*udpSlot pool []*udpSlot
reserver Reserver
resetter Resetter backing []byte
} }
// NewUDPCoalescer wraps w. The caller is responsible for only constructing // NewUDPCoalescer wraps w. The caller is responsible for only constructing
@@ -71,14 +74,13 @@ type UDPCoalescer struct {
// the kernel may reject GSO_UDP_L4 writes. If w does not implement // the kernel may reject GSO_UDP_L4 writes. If w does not implement
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to // tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
// plain Writes — same defensive shape as the TCP coalescer. // plain Writes — same defensive shape as the TCP coalescer.
func NewUDPCoalescer(w io.Writer, reserver Reserver, resetter Resetter) *UDPCoalescer { func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
c := &UDPCoalescer{ c := &UDPCoalescer{
plainW: w, plainW: w,
slots: make([]*udpSlot, 0, initialSlots), slots: make([]*udpSlot, 0, initialSlots),
openSlots: make(map[flowKey]*udpSlot, initialSlots), openSlots: make(map[flowKey]*udpSlot, initialSlots),
pool: make([]*udpSlot, 0, initialSlots), pool: make([]*udpSlot, 0, initialSlots),
reserver: reserver, backing: make([]byte, 0, initialSlots*udpCoalesceBufSize),
resetter: resetter,
} }
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok { if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
c.gsoW = gw c.gsoW = gw
@@ -124,7 +126,7 @@ func parseUDP(pkt []byte) (parsedUDP, bool) {
} }
func (c *UDPCoalescer) Reserve(sz int) []byte { func (c *UDPCoalescer) Reserve(sz int) []byte {
return c.reserver(sz) return reserveFromBacking(&c.backing, sz)
} }
// Commit borrows pkt. The caller must keep pkt valid until the next Flush. // Commit borrows pkt. The caller must keep pkt valid until the next Flush.
@@ -149,17 +151,6 @@ func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
c.addPassthrough(pkt) c.addPassthrough(pkt)
return nil return nil
} }
// A zero-length UDP datagram (UDP `length` == 8) is legal and must still
// reach the TUN, but it can't be coalesced: a GSO slot would store an
// empty payload iovec and the kernel has nothing to segment. Seal any
// open chain for this flow (so a later, non-empty datagram seeds fresh
// *after* this one and per-flow arrival order is preserved) and deliver
// it as a plain single datagram.
if info.payLen == 0 {
delete(c.openSlots, info.fk)
c.addPassthrough(pkt)
return nil
}
if open := c.openSlots[info.fk]; open != nil { if open := c.openSlots[info.fk]; open != nil {
if c.canAppend(open, pkt, info) { if c.canAppend(open, pkt, info) {
c.appendPayload(open, pkt, info) c.appendPayload(open, pkt, info)
@@ -175,19 +166,7 @@ func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
return nil return nil
} }
// Flush drains every queued slot and calls the configured Resetter.
func (c *UDPCoalescer) Flush() error { func (c *UDPCoalescer) Flush() error {
first := c.drain()
if c.resetter != nil {
c.resetter()
}
return first
}
// drain emits every queued slot in arrival order and clears the slot state.
// It does NOT reset the arena: borrowed payload slices stay valid until the
// arena's owner recycles it.
func (c *UDPCoalescer) drain() error {
var first error var first error
for _, s := range c.slots { for _, s := range c.slots {
var err error var err error
@@ -204,6 +183,7 @@ func (c *UDPCoalescer) drain() error {
clear(c.slots) clear(c.slots)
c.slots = c.slots[:0] c.slots = c.slots[:0]
clear(c.openSlots) clear(c.openSlots)
c.backing = c.backing[:0]
return first return first
} }
@@ -265,6 +245,8 @@ func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen]) s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
s.numSeg++ s.numSeg++
s.totalPay += info.payLen s.totalPay += info.payLen
// Merge IP-level CE marks into the seed (same trick TCP coalescer uses).
mergeECNIntoSeed(s.hdrBuf[:s.ipHdrLen], pkt[:s.ipHdrLen], s.isV6)
if info.payLen < s.gsoSize { if info.payLen < s.gsoSize {
// Last-segment-can-be-shorter: this seals the chain. // Last-segment-can-be-shorter: this seals the chain.
s.sealed = true s.sealed = true
@@ -335,9 +317,8 @@ func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on // udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
// every field that must be identical across coalesced segments. Length // every field that must be identical across coalesced segments. Length
// fields are masked out (flushSlot rewrites them), but the IP-level ECN // fields and the ECN bits in IP TOS/TC are masked out — appendPayload
// codepoint is compared (via ipHeadersMatch) so segments with differing ECN // merges CE into the seed; flushSlot rewrites lengths.
// don't coalesce, matching kernel GRO.
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool { func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
if len(a) != len(b) { if len(a) != len(b) {
return false return false
+28 -117
View File
@@ -60,8 +60,7 @@ func buildUDPv6(sport, dport uint16, payload []byte) []byte {
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) { func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: false} w := &fakeTunWriter{gsoEnabled: false}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 100)) pkt := buildUDPv4(1000, 53, make([]byte, 100))
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -79,8 +78,7 @@ func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) { func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
// ICMP packet // ICMP packet
pkt := make([]byte, 28) pkt := make([]byte, 28)
pkt[0] = 0x45 pkt[0] = 0x45
@@ -101,8 +99,7 @@ func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) { func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 800)) pkt := buildUDPv4(1000, 53, make([]byte, 800))
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -119,8 +116,7 @@ func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) { func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
@@ -160,8 +156,7 @@ func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
// Last segment may be shorter, sealing the chain. // Last segment may be shorter, sealing the chain.
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) { func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
full := make([]byte, 1200) full := make([]byte, 1200)
tail := make([]byte, 600) tail := make([]byte, 600)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
@@ -194,8 +189,7 @@ func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
// A larger-than-gsoSize packet cannot extend the slot — it reseeds. // A larger-than-gsoSize packet cannot extend the slot — it reseeds.
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) { func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil { if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -213,8 +207,7 @@ func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
// Different 5-tuples must not coalesce. // Different 5-tuples must not coalesce.
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) { func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800) pay := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -245,8 +238,7 @@ func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
// Caps at udpCoalesceMaxSegs. // Caps at udpCoalesceMaxSegs.
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) { func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 100) pay := make([]byte, 100)
for i := 0; i < udpCoalesceMaxSegs+5; i++ { for i := 0; i < udpCoalesceMaxSegs+5; i++ {
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
@@ -269,46 +261,39 @@ func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
} }
} }
// Differing IP ECN codepoints must not coalesce: udpHeadersMatch compares // CE marks on appended segments must be merged into the seed's IP TOS.
// the full ToS byte (matching kernel GRO). A CE-marked datagram mid-run func TestUDPCoalescerMergesCEMark(t *testing.T) {
// seals the Not-ECT chain and seeds a fresh superpacket that keeps CE; the
// trailing Not-ECT datagram seeds another.
func TestUDPCoalescerDifferingECNReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800) pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00 (Not-ECT) pkt0 := buildUDPv4(1000, 53, pay) // ECN=00
pkt1 := buildUDPv4(1000, 53, pay) pkt1 := buildUDPv4(1000, 53, pay)
pkt1[1] = 0x03 // CE pkt1[1] = 0x03 // CE
pkt2 := buildUDPv4(1000, 53, pay) // ECN=00 again pkt2 := buildUDPv4(1000, 53, pay)
for _, p := range [][]byte{pkt0, pkt1, pkt2} { if err := c.Commit(pkt0); err != nil {
if err := c.Commit(p); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if err := c.Commit(pkt1); err != nil {
t.Fatal(err)
}
if err := c.Commit(pkt2); err != nil {
t.Fatal(err)
} }
if err := c.Flush(); err != nil { if err := c.Flush(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(w.gsoWrites) != 3 { if len(w.gsoWrites) != 1 {
t.Fatalf("want 3 separate seeds (differing ECN), got %d (plain=%d)", len(w.gsoWrites), len(w.writes)) t.Fatalf("want 1 merged gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
}
wantECN := []byte{0x00, 0x03, 0x00}
for i, g := range w.gsoWrites {
if len(g.pays) != 1 {
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
}
if got := g.hdr[1] & 0x03; got != wantECN[i] {
t.Errorf("gso %d ECN=%#x want %#x", i, got, wantECN[i])
} }
if w.gsoWrites[0].hdr[1]&0x03 != 0x03 {
t.Errorf("CE not merged into seed (tos=%#x)", w.gsoWrites[0].hdr[1])
} }
} }
// IPv6 path: same flow, equal-sized → coalesced. // IPv6 path: same flow, equal-sized → coalesced.
func TestUDPCoalescerIPv6Coalesces(t *testing.T) { func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 1200) pay := make([]byte, 1200)
for i := 0; i < 3; i++ { for i := 0; i < 3; i++ {
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil { if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
@@ -341,11 +326,10 @@ func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
} }
} }
// DSCP differences must reseed: udpHeadersMatch compares the full ToS byte. // DSCP differences must reseed (headers don't match outside ECN).
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) { func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pay := make([]byte, 800) pay := make([]byte, 800)
pkt0 := buildUDPv4(1000, 53, pay) pkt0 := buildUDPv4(1000, 53, pay)
pkt1 := buildUDPv4(1000, 53, pay) pkt1 := buildUDPv4(1000, 53, pay)
@@ -367,8 +351,7 @@ func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
// Fragmented IPv4 must not be coalesced. // Fragmented IPv4 must not be coalesced.
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) { func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 200)) pkt := buildUDPv4(1000, 53, make([]byte, 200))
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1 binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
@@ -382,82 +365,10 @@ func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
} }
} }
// A zero-length UDP datagram (UDP length == 8, no payload) is legal and
// must be delivered as a plain single datagram — never coalesced. Seeding
// it into a GSO slot stores an empty payload iovec that panics WriteGSO
// (index-out-of-range on &pay[0]); this is a remote DoS if we ever let it
// reach the GSO path. Regression: must not panic and must be written.
func TestUDPCoalescerZeroLengthPayloadPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("zero-length UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
if len(w.writes[0]) != len(pkt) {
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
}
}
// IPv6 zero-length UDP datagram: same passthrough contract as v4.
func TestUDPCoalescerZeroLengthPayloadIPv6PassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv6(1000, 53, nil) // UDP length 8, zero payload
if err := c.Commit(pkt); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
t.Fatalf("zero-length IPv6 UDP must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
}
if len(w.writes[0]) != len(pkt) {
t.Errorf("delivered %d bytes, want the whole %d-byte datagram", len(w.writes[0]), len(pkt))
}
}
// A zero-length datagram arriving mid-flow must seal the open chain so the
// datagram after it seeds a fresh superpacket *after* the empty one on the
// wire — per-flow arrival order (full, empty, full) must be preserved.
func TestUDPCoalescerZeroLengthMidFlowSealsAndPreservesOrder(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
full := make([]byte, 800)
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, nil)); err != nil { // zero-length
t.Fatal(err)
}
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
t.Fatal(err)
}
if err := c.Flush(); err != nil {
t.Fatal(err)
}
// The empty datagram sealed the first slot, so the trailing full packet
// can't join it: two single-segment superpackets bracket one plain write.
if len(w.gsoWrites) != 2 || len(w.writes) != 1 {
t.Fatalf("want 2 gso writes + 1 plain, got gso=%d plain=%d", len(w.gsoWrites), len(w.writes))
}
}
// IPv4 with options is not admissible (we require IHL=5). // IPv4 with options is not admissible (we require IHL=5).
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) { func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
w := &fakeTunWriter{gsoEnabled: true} w := &fakeTunWriter{gsoEnabled: true}
arena := NewArena(0) c := NewUDPCoalescer(w)
c := NewUDPCoalescer(w, arena.Reserve, arena.Reset)
pkt := buildUDPv4(1000, 53, make([]byte, 200)) pkt := buildUDPv4(1000, 53, make([]byte, 200))
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options) pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
if err := c.Commit(pkt); err != nil { if err := c.Commit(pkt); err != nil {
+3 -7
View File
@@ -18,11 +18,7 @@ type Device interface {
Networks() []netip.Prefix Networks() []netip.Prefix
Name() string Name() string
RoutesFor(netip.Addr) routing.Gateways RoutesFor(netip.Addr) routing.Gateways
// Queues returns the device's packet queues, opening additional ones as SupportsMultiqueue() bool
// needed until there are n. Platforms without multiqueue support return NewMultiQueueReader() error
// their single queue regardless of n, so callers must size reader loops Readers() []tio.Queue
// to len(result), not n; implementations never return more than n. An
// error means a queue that should have opened could not; the caller owns
// cleanup via Close. Called once, during interface activation.
Queues(n int) ([]tio.Queue, error)
} }
+11 -2
View File
@@ -3,6 +3,7 @@
package overlaytest package overlaytest
import ( import (
"errors"
"net/netip" "net/netip"
"github.com/slackhq/nebula/overlay/tio" "github.com/slackhq/nebula/overlay/tio"
@@ -38,8 +39,16 @@ func (NoopTun) Write([]byte) (int, error) {
return 0, nil return 0, nil
} }
func (NoopTun) Queues(int) ([]tio.Queue, error) { func (NoopTun) SupportsMultiqueue() bool {
return []tio.Queue{NoopTun{}}, nil return false
}
func (NoopTun) NewMultiQueueReader() error {
return errors.New("unsupported")
}
func (NoopTun) Readers() []tio.Queue {
return []tio.Queue{NoopTun{}}
} }
func (NoopTun) Close() error { func (NoopTun) Close() error {
-45
View File
@@ -1,45 +0,0 @@
//go:build linux && !android
// +build linux,!android
package tio
import (
"os"
"golang.org/x/sys/unix"
)
// blockOn parks the calling goroutine until fd is ready (events is POLLIN for
// reads, POLLOUT for writes) or shutdownFd signals teardown. It builds the
// pollfd array on the stack every call, so concurrent callers on the same
// Queue never share Revents storage.
//
// Returns os.ErrClosed when shutdown was signaled (POLLIN on shutdownFd)
// or either fd reported a problem condition (POLLHUP|POLLNVAL|POLLERR).
func blockOn(fd, shutdownFd int32, events int16) error {
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
pfds := [2]unix.PollFd{
{Fd: fd, Events: events},
{Fd: shutdownFd, Events: unix.POLLIN},
}
var err error
for {
_, err = unix.Poll(pfds[:], -1)
if err != unix.EINTR {
break
}
}
tunEvents := pfds[0].Revents
shutdownEvents := pfds[1].Revents
// Check err before trusting the potentially bogus bits we just got.
if err != nil {
return err
}
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
return os.ErrClosed
}
if tunEvents&problemFlags != 0 {
return os.ErrClosed
}
return nil
}
+1 -21
View File
@@ -1,13 +1,9 @@
//go:build linux && !android
// +build linux,!android
package tio package tio
import ( import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"fmt" "fmt"
"sync/atomic"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
@@ -21,7 +17,6 @@ type offloadQueueSet struct {
// with the kernel. Queues created by Add inherit this and surface it // with the kernel. Queues created by Add inherit this and surface it
// via Offload.USOSupported so coalescers can gate USO emission. // via Offload.USOSupported so coalescers can gate USO emission.
usoEnabled bool usoEnabled bool
closed atomic.Bool
} }
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do // NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
@@ -67,33 +62,18 @@ func (c *offloadQueueSet) wakeForShutdown() error {
} }
func (c *offloadQueueSet) Close() error { func (c *offloadQueueSet) Close() error {
if c.closed.Swap(true) {
return nil
}
errs := []error{} errs := []error{}
// Signal all readers blocked in poll to wake up and exit. They observe // Signal all readers blocked in poll to wake up and exit
// POLLIN on the shutdown eventfd and return os.ErrClosed.
if err := c.wakeForShutdown(); err != nil { if err := c.wakeForShutdown(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
// Close the per-queue tun fds; this also unblocks any in-flight reads.
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
// to this container.
for _, x := range c.pq { for _, x := range c.pq {
if err := x.Close(); err != nil { if err := x.Close(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
} }
// Close the shutdown eventfd last: every reader's pollfd set references
// it, so it must outlive the wake + per-queue teardown above.
if err := unix.Close(c.shutdownFd); err != nil {
errs = append(errs, err)
}
c.shutdownFd = -1
return errors.Join(errs...) return errors.Join(errs...)
} }
-21
View File
@@ -1,13 +1,9 @@
//go:build linux && !android
// +build linux,!android
package tio package tio
import ( import (
"encoding/binary" "encoding/binary"
"errors" "errors"
"fmt" "fmt"
"sync/atomic"
"golang.org/x/sys/unix" "golang.org/x/sys/unix"
) )
@@ -17,7 +13,6 @@ type pollQueueSet struct {
// pqi is exactly the same as pq, but stored as the interface type // pqi is exactly the same as pq, but stored as the interface type
pqi []Queue pqi []Queue
shutdownFd int shutdownFd int
closed atomic.Bool
} }
func NewPollQueueSet() (QueueSet, error) { func NewPollQueueSet() (QueueSet, error) {
@@ -58,33 +53,17 @@ func (c *pollQueueSet) wakeForShutdown() error {
} }
func (c *pollQueueSet) Close() error { func (c *pollQueueSet) Close() error {
if c.closed.Swap(true) {
return nil
}
errs := []error{} errs := []error{}
// Wake any reader blocked in poll so it observes POLLIN on the shutdown
// eventfd and returns os.ErrClosed.
if err := c.wakeForShutdown(); err != nil { if err := c.wakeForShutdown(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
// Close the per-queue tun fds; this also unblocks any in-flight reads.
// The per-queue Close deliberately leaves shutdownFd alone - it belongs
// to this container.
for _, x := range c.pq { for _, x := range c.pq {
if err := x.Close(); err != nil { if err := x.Close(); err != nil {
errs = append(errs, err) errs = append(errs, err)
} }
} }
// Close the shutdown eventfd last: every reader's pollfd set references
// it, so it must outlive the wake + per-queue teardown above.
if err := unix.Close(c.shutdownFd); err != nil {
errs = append(errs, err)
}
c.shutdownFd = -1
return errors.Join(errs...) return errors.Join(errs...)
} }
-50
View File
@@ -1,50 +0,0 @@
package tio
import "io"
// singleQueue adapts a legacy one-datagram-per-Read source into a Queue.
// Read fills a private scratch buffer and returns exactly one Packet whose
// Bytes borrow from that buffer, valid only until the next Read, per the
// Queue contract. Single-reader like every Queue; Write is exactly as safe
// for concurrent use as the underlying source's Write.
type singleQueue struct {
rw io.ReadWriter
closer io.Closer // nil: Close is a no-op (the source is shared and owned elsewhere)
buf []byte
ret [1]Packet
}
// NewSingleQueue wraps a one-datagram-per-Read ReadWriteCloser (a legacy tun
// device) into a Queue. bufSize is the per-queue read scratch size and must
// be at least the largest datagram the source can return. Close closes rwc.
func NewSingleQueue(rwc io.ReadWriteCloser, bufSize int) Queue {
return &singleQueue{rw: rwc, closer: rwc, buf: make([]byte, bufSize)}
}
// NewSingleQueueNoClose is NewSingleQueue for a source owned by someone else,
// e.g. several queues sharing one device. Close on the returned Queue is a
// no-op so one queue can't tear the shared source out from under its
// siblings; the owner remains responsible for closing the source itself.
func NewSingleQueueNoClose(rw io.ReadWriter, bufSize int) Queue {
return &singleQueue{rw: rw, buf: make([]byte, bufSize)}
}
func (q *singleQueue) Read() ([]Packet, error) {
n, err := q.rw.Read(q.buf)
if err != nil {
return nil, err
}
q.ret[0] = Packet{Bytes: q.buf[:n]}
return q.ret[:], nil
}
func (q *singleQueue) Write(p []byte) (int, error) {
return q.rw.Write(p)
}
func (q *singleQueue) Close() error {
if q.closer == nil {
return nil
}
return q.closer.Close()
}
+4 -6
View File
@@ -26,9 +26,8 @@ type Capabilities struct {
USO bool USO bool
} }
// Queue is a readable/writable Poll queue. Concurrency contract: a single // Queue is a readable/writable Poll queue. One Queue is driven by a single
// read goroutine drives Read; plain Write is safe for concurrent callers; // read goroutine plus a single writer (see Write below).
// WriteGSO (on Queues that implement GSOWriter) is single-writer per queue.
type Queue interface { type Queue interface {
io.Closer io.Closer
@@ -38,12 +37,11 @@ type Queue interface {
// or copy each slice before the next call. A Packet may carry a // or copy each slice before the next call. A Packet may carry a
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is // GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is
// true the caller must segment Bytes before treating it as a single // true the caller must segment Bytes before treating it as a single
// IP datagram. Single-reader only: not safe for concurrent Reads (it // IP datagram. Not safe for concurrent Reads.
// reuses per-queue rx scratch each call).
Read() ([]Packet, error) Read() ([]Packet, error)
// Write emits a single packet on the plaintext (outside→inside) // Write emits a single packet on the plaintext (outside→inside)
// delivery path. Safe for concurrent use. // delivery path. Not safe for concurrent Writes.
Write(p []byte) (int, error) Write(p []byte) (int, error)
} }
+78 -31
View File
@@ -1,6 +1,3 @@
//go:build linux && !android
// +build linux,!android
package tio package tio
import ( import (
@@ -8,6 +5,7 @@ import (
"io" "io"
"log/slog" "log/slog"
"os" "os"
"sync"
"sync/atomic" "sync/atomic"
"syscall" "syscall"
"unsafe" "unsafe"
@@ -28,12 +26,9 @@ const tunRxBufSize = 64 * 1024
// tunRxBufCap is the total size we allocate for the per-reader rx // tunRxBufCap is the total size we allocate for the per-reader rx
// buffer. With reads landing directly in rxBuf, each drain iteration // buffer. With reads landing directly in rxBuf, each drain iteration
// consumes up to tunRxBufSize of headroom for the kernel-supplied bytes. // consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
// Sized to eight such iterations so a single poll wake can drain several // Sized to two such iterations so the initial blocking read plus one
// TSO/USO superpackets under bulk load, amortizing the wake and giving // drain read both fit without partial-drop.
// the sendmmsg planner longer same-destination runs. Hold latency stays const tunRxBufCap = tunRxBufSize * 2
// bounded because listenIn flushes its send batch incrementally rather
// than only at end-of-drain.
const tunRxBufCap = tunRxBufSize * 8
// tunDrainCap caps how many packets a single Read will accumulate via // tunDrainCap caps how many packets a single Read will accumulate via
// the post-wake drain loop. Sized to soak up a burst of small ACKs while // the post-wake drain loop. Sized to soak up a burst of small ACKs while
@@ -64,6 +59,13 @@ var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
type Offload struct { type Offload struct {
fd int fd int
shutdownFd int shutdownFd int
readPoll [2]unix.PollFd
writePoll [2]unix.PollFd
// writeLock serializes blockOnWrite's read+clear of writePoll[*].Revents.
// Any goroutine that calls Write may end up parked in poll(2); without
// the lock concurrent waiters could race the Revents reset and lose
// events.
writeLock sync.Mutex
closed atomic.Bool closed atomic.Bool
rxBuf []byte // backing store for kernel-handed packets read this drain rxBuf []byte // backing store for kernel-handed packets read this drain
rxOff int // cursor into rxBuf for the current Read drain rxOff int // cursor into rxBuf for the current Read drain
@@ -104,6 +106,15 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
shutdownFd: shutdownFd, shutdownFd: shutdownFd,
usoEnabled: usoEnabled, usoEnabled: usoEnabled,
closed: atomic.Bool{}, closed: atomic.Bool{},
readPoll: [2]unix.PollFd{
{Fd: int32(fd), Events: unix.POLLIN},
{Fd: int32(shutdownFd), Events: unix.POLLIN},
},
writePoll: [2]unix.PollFd{
{Fd: int32(fd), Events: unix.POLLOUT},
{Fd: int32(shutdownFd), Events: unix.POLLIN},
},
writeLock: sync.Mutex{},
rxBuf: make([]byte, tunRxBufCap), rxBuf: make([]byte, tunRxBufCap),
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs), gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
@@ -121,11 +132,57 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
} }
func (r *Offload) blockOnRead() error { func (r *Offload) blockOnRead() error {
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLIN) const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
var err error
for {
_, err = unix.Poll(r.readPoll[:], -1)
if err != unix.EINTR {
break
}
}
//always reset these!
tunEvents := r.readPoll[0].Revents
shutdownEvents := r.readPoll[1].Revents
r.readPoll[0].Revents = 0
r.readPoll[1].Revents = 0
//do the err check before trusting the potentially bogus bits we just got
if err != nil {
return err
}
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
return os.ErrClosed
} else if tunEvents&problemFlags != 0 {
return os.ErrClosed
}
return nil
} }
func (r *Offload) blockOnWrite() error { func (r *Offload) blockOnWrite() error {
return blockOn(int32(r.fd), int32(r.shutdownFd), unix.POLLOUT) const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
var err error
for {
_, err = unix.Poll(r.writePoll[:], -1)
if err != unix.EINTR {
break
}
}
//always reset these!
r.writeLock.Lock()
tunEvents := r.writePoll[0].Revents
shutdownEvents := r.writePoll[1].Revents
r.writePoll[0].Revents = 0
r.writePoll[1].Revents = 0
r.writeLock.Unlock()
//do the err check before trusting the potentially bogus bits we just got
if err != nil {
return err
}
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
return os.ErrClosed
} else if tunEvents&problemFlags != 0 {
return os.ErrClosed
}
return nil
} }
// readPacket issues a single readv(2) splitting the virtio_net_hdr off // readPacket issues a single readv(2) splitting the virtio_net_hdr off
@@ -379,22 +436,10 @@ func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto
r.gsoIovs[1].SetLen(len(hdr)) r.gsoIovs[1].SetLen(len(hdr))
r.gsoIovs[2].Base = &transportHdr[0] r.gsoIovs[2].Base = &transportHdr[0]
r.gsoIovs[2].SetLen(len(transportHdr)) r.gsoIovs[2].SetLen(len(transportHdr))
// Defense in depth: an empty payload fragment can't be a valid GSO for i, p := range pays {
// segment and &p[0] would panic on it. Callers route zero-length r.gsoIovs[3+i].Base = &p[0]
// datagrams through the plain path (see UDPCoalescer.commitParsed), so r.gsoIovs[3+i].SetLen(len(p))
// this should never fire, but skip empties rather than index into one.
// `n` tracks where the next payload iovec lands, since skips make it
// drift from 3+i.
n := 3
for _, p := range pays {
if len(p) == 0 {
continue
} }
r.gsoIovs[n].Base = &p[0]
r.gsoIovs[n].SetLen(len(p))
n++
}
r.gsoIovs = r.gsoIovs[:n]
_, err := r.rawWrite(r.gsoIovs) _, err := r.rawWrite(r.gsoIovs)
return err return err
@@ -406,9 +451,11 @@ func (r *Offload) Close() error {
} }
//shutdownFd is owned by the container, so we should not close it //shutdownFd is owned by the container, so we should not close it
// Close the underlying fd but do NOT null r.fd: a reader may still be var err error
// loading it in readOne, and mutating the field would race that load. if r.fd >= 0 {
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's err = unix.Close(r.fd)
// ppoll first). closed.Swap already guarantees we only close once. r.fd = -1
return unix.Close(r.fd) }
return err
} }
+64 -17
View File
@@ -1,6 +1,3 @@
//go:build linux && !android
// +build linux,!android
package tio package tio
import ( import (
@@ -18,26 +15,32 @@ const tunReadBufSize = 65535
type Poll struct { type Poll struct {
fd int fd int
shutdownFd int
readPoll [2]unix.PollFd
writePoll [2]unix.PollFd
closed atomic.Bool closed atomic.Bool
readBuf []byte readBuf []byte
batchRet [1]Packet batchRet [1]Packet
} }
// newPoll wraps an existing tun fd. On failure it does NOT close fd: the
// caller owns fd and is the sole closer (see pollQueueSet.Add callers in
// overlay/tun_linux.go, which unix.Close on Add error). This matches the
// newOffload convention and keeps closes at exactly one on every path.
func newPoll(fd int, shutdownFd int) (*Poll, error) { func newPoll(fd int, shutdownFd int) (*Poll, error) {
if err := unix.SetNonblock(fd, true); err != nil { if err := unix.SetNonblock(fd, true); err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err) return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
} }
out := &Poll{ out := &Poll{
fd: fd, fd: fd,
shutdownFd: shutdownFd,
readBuf: make([]byte, tunReadBufSize), readBuf: make([]byte, tunReadBufSize),
readPoll: [2]unix.PollFd{
{Fd: int32(fd), Events: unix.POLLIN},
{Fd: int32(shutdownFd), Events: unix.POLLIN},
},
writePoll: [2]unix.PollFd{
{Fd: int32(fd), Events: unix.POLLOUT},
{Fd: int32(shutdownFd), Events: unix.POLLIN},
},
} }
return out, nil return out, nil
} }
@@ -45,11 +48,53 @@ func newPoll(fd int, shutdownFd int) (*Poll, error) {
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled. // blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
// Returns os.ErrClosed if Close was called. // Returns os.ErrClosed if Close was called.
func (t *Poll) blockOnRead() error { func (t *Poll) blockOnRead() error {
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLIN) const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
var err error
for {
_, err = unix.Poll(t.readPoll[:], -1)
if err != unix.EINTR {
break
}
}
tunEvents := t.readPoll[0].Revents
shutdownEvents := t.readPoll[1].Revents
t.readPoll[0].Revents = 0
t.readPoll[1].Revents = 0
if err != nil {
return err
}
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
return os.ErrClosed
}
if tunEvents&problemFlags != 0 {
return os.ErrClosed
}
return nil
} }
func (t *Poll) blockOnWrite() error { func (t *Poll) blockOnWrite() error {
return blockOn(int32(t.fd), int32(t.shutdownFd), unix.POLLOUT) const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
var err error
for {
_, err = unix.Poll(t.writePoll[:], -1)
if err != unix.EINTR {
break
}
}
tunEvents := t.writePoll[0].Revents
shutdownEvents := t.writePoll[1].Revents
t.writePoll[0].Revents = 0
t.writePoll[1].Revents = 0
if err != nil {
return err
}
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
return os.ErrClosed
}
if tunEvents&problemFlags != 0 {
return os.ErrClosed
}
return nil
} }
func (t *Poll) Read() ([]Packet, error) { func (t *Poll) Read() ([]Packet, error) {
@@ -82,7 +127,7 @@ func (t *Poll) readOne(to []byte) (int, error) {
} }
} }
// Write is safe for concurrent use // Write is only valid for single threaded use
func (t *Poll) Write(from []byte) (int, error) { func (t *Poll) Write(from []byte) (int, error) {
for { for {
n, errno := unix.Write(t.fd, from) n, errno := unix.Write(t.fd, from)
@@ -109,9 +154,11 @@ func (t *Poll) Close() error {
return nil return nil
} }
//shutdownFd is owned by the container, so we should not close it //shutdownFd is owned by the container, so we should not close it
// Close the underlying fd but do NOT null t.fd: a reader may still be var err error
// loading it in readOne, and mutating the field would race that load. if t.fd >= 0 {
// It gets EBADF -> os.ErrClosed (or wakes via the shutdown eventfd's err = unix.Close(t.fd)
// ppoll first). closed.Swap already guarantees we only close once. t.fd = -1
return unix.Close(t.fd) }
return err
} }
-145
View File
@@ -70,97 +70,6 @@ func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
} }
} }
// TestPoll_ConcurrentWrite_NoRace hammers a single Poll queue from two writer
// goroutines while a reader drains the other end of the pipe. The writers
// overflow the pipe buffer, so both repeatedly park in blockOnWrite at the same
// time — the exact scenario that raced on the old shared writePoll member
// array. Run under -race; a shared-array regression trips the detector here.
func TestPoll_ConcurrentWrite_NoRace(t *testing.T) {
var fds [2]int
require.NoError(t, unix.Pipe2(fds[:], unix.O_CLOEXEC))
readFd, writeFd := fds[0], fds[1]
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
require.NoError(t, err)
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
p, err := newPoll(writeFd, shutdownFd)
require.NoError(t, err)
const writers = 2
const perWriter = 4000
payload := make([]byte, 100)
total := writers * perWriter * len(payload)
// Reader: drain the read end (blocking) until every writer's bytes are
// consumed, so the writers keep making progress rather than wedging on a
// permanently full pipe.
readDone := make(chan struct{})
go func() {
defer close(readDone)
buf := make([]byte, 4096)
got := 0
for got < total {
n, rerr := unix.Read(readFd, buf)
got += n
if rerr != nil {
if rerr == unix.EINTR {
continue
}
return
}
if n == 0 { // EOF
return
}
}
}()
var wg sync.WaitGroup
for w := 0; w < writers; w++ {
wg.Add(1)
go func() {
defer wg.Done()
for i := 0; i < perWriter; i++ {
if _, werr := p.Write(payload); werr != nil {
t.Errorf("write: %v", werr)
return
}
}
}()
}
wg.Wait()
select {
case <-readDone:
case <-time.After(10 * time.Second):
t.Fatal("reader did not drain")
}
require.NoError(t, p.Close())
_ = unix.Close(readFd)
}
// TestPoll_NewPoll_DoesNotCloseFdOnFailure pins the ownership rule: when
// newPoll fails, it must leave fd open so the caller (pollQueueSet.Add's
// callers in tun_linux.go) is the sole closer. If newPoll also closed fd,
// the poll path would double-close on Add error. We force the failure with
// an O_PATH descriptor: fcntl(F_SETFL) — which SetNonblock performs — is not
// permitted on O_PATH fds and fails with EBADF, while the fd itself stays
// open so we can observe that newPoll left it alone.
func TestPoll_NewPoll_DoesNotCloseFdOnFailure(t *testing.T) {
fd, err := unix.Open("/", unix.O_PATH|unix.O_CLOEXEC, 0)
require.NoError(t, err)
t.Cleanup(func() { _ = unix.Close(fd) })
p, err := newPoll(fd, 1)
require.Error(t, err, "SetNonblock on an O_PATH fd should fail")
require.Nil(t, p)
// If newPoll had closed fd, F_GETFD would report it closed. It staying
// open proves newPoll left the fd for the caller to close exactly once.
require.True(t, fdOpen(t, fd), "newPoll must not close fd on failure; caller is the sole closer")
}
func TestPoll_Close_Idempotent(t *testing.T) { func TestPoll_Close_Idempotent(t *testing.T) {
tf, err := newPoll(newReadPipe(t), 1) tf, err := newPoll(newReadPipe(t), 1)
require.NoError(t, err) require.NoError(t, err)
@@ -171,57 +80,3 @@ func TestPoll_Close_Idempotent(t *testing.T) {
t.Fatalf("second Close should be a no-op, got %v", err) t.Fatalf("second Close should be a no-op, got %v", err)
} }
} }
// fdOpen reports whether fd currently refers to an open file description.
// A closed (or never-allocated) fd makes F_GETFD fail with EBADF.
func fdOpen(t *testing.T, fd int) bool {
t.Helper()
_, err := unix.FcntlInt(uintptr(fd), unix.F_GETFD, 0)
if err == nil {
return true
}
if errors.Is(err, unix.EBADF) {
return false
}
t.Fatalf("unexpected fcntl(F_GETFD) error on fd %d: %v", fd, err)
return false
}
// TestPollQueueSet_Close_ClosesShutdownFd is the regression test for the
// leaked shutdown eventfd: the container that owns shutdownFd must close it in
// Close, and a second Close must be a safe no-op.
func TestPollQueueSet_Close_ClosesShutdownFd(t *testing.T) {
qs, err := NewPollQueueSet()
require.NoError(t, err)
c, ok := qs.(*pollQueueSet)
require.True(t, ok)
require.NoError(t, qs.Add(newReadPipe(t)))
shutdownFd := c.shutdownFd
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
require.NoError(t, qs.Close())
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
require.NoError(t, qs.Close())
}
// TestOffloadQueueSet_Close_ClosesShutdownFd mirrors the poll regression test
// for the GSO/offload queueset.
func TestOffloadQueueSet_Close_ClosesShutdownFd(t *testing.T) {
qs, err := NewOffloadQueueSet(false)
require.NoError(t, err)
c, ok := qs.(*offloadQueueSet)
require.True(t, ok)
require.NoError(t, qs.Add(newReadPipe(t)))
shutdownFd := c.shutdownFd
require.True(t, fdOpen(t, shutdownFd), "shutdown eventfd should be open before Close")
require.NoError(t, qs.Close())
require.False(t, fdOpen(t, shutdownFd), "shutdown eventfd should be closed after Close")
// Second Close must not touch fds (shutdownFd is now -1) and must return nil.
require.NoError(t, qs.Close())
}
-32
View File
@@ -640,38 +640,6 @@ func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
} }
} }
// TestWriteGSOSkipsEmptyPayloads is the defense-in-depth guard for the
// zero-length UDP DoS: a payload fragment of length zero would make &p[0]
// panic (index-out-of-range) when building the iovec array. WriteGSO must
// skip empties instead. We write to /dev/null so the writev always succeeds
// synchronously; the point is simply that neither call panics.
func TestWriteGSOSkipsEmptyPayloads(t *testing.T) {
fd, err := unix.Open("/dev/null", os.O_WRONLY, 0)
if err != nil {
t.Fatalf("open /dev/null: %v", err)
}
t.Cleanup(func() { _ = unix.Close(fd) })
o := &Offload{fd: fd, gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs)}
o.gsoIovs[0].Base = &o.gsoHdrBuf[0]
o.gsoIovs[0].SetLen(virtio.Size)
ipHdr := make([]byte, 20)
ipHdr[0] = 0x45 // IPv4, IHL 5
udpHdr := make([]byte, 8)
// Sole payload empty: exercises the all-empty skip (n stays at 3).
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{{}}, GSOProtoUDP); err != nil {
t.Fatalf("WriteGSO with a single empty payload: %v", err)
}
// Empty mixed with a real fragment: exercises the index-drift skip so a
// later non-empty payload still lands in the right iovec slot.
real := make([]byte, 1200)
if err := o.WriteGSO(ipHdr, udpHdr, [][]byte{real, {}}, GSOProtoUDP); err != nil {
t.Fatalf("WriteGSO with a trailing empty payload: %v", err)
}
}
// buildTSOv6 builds a synthetic IPv6/TCP TSO superpacket with payLen bytes // buildTSOv6 builds a synthetic IPv6/TCP TSO superpacket with payLen bytes
// of payload, segmented at gso. Returns the packet bytes only; the // of payload, segmented at gso. Returns the packet bytes only; the
// virtio_net_hdr is the caller's responsibility. // virtio_net_hdr is the caller's responsibility.
-3
View File
@@ -1,3 +0,0 @@
//go:build !linux || android
package virtio
+21 -53
View File
@@ -27,13 +27,6 @@ const (
tcpHeaderMaxLen = 60 // data-offset=15, max options tcpHeaderMaxLen = 60 // data-offset=15, max options
) )
// maxSegHdrLen bounds the L3+L4 header we snapshot before stamping each
// segment. The largest header the segmenter supports is IPv4 (max IHL 60)
// plus TCP (max data-offset 60) = 120 bytes; the array is sized to that
// worst case so the snapshot lives on the stack with no per-call heap
// allocation.
const maxSegHdrLen = ipv4HeaderMaxLen + tcpHeaderMaxLen // 120
// Byte offsets inside an IPv4 header. // Byte offsets inside an IPv4 header.
const ( const (
ipv4TotalLenOff = 2 ipv4TotalLenOff = 2
@@ -143,7 +136,7 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
if hdr.HdrLen < hdr.CsumStart { if hdr.HdrLen < hdr.CsumStart {
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart) return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
} }
cSumAt := int(hdr.CsumStart + hdr.CsumOffset) cSumAt := int(hdr.CsumStart + hdr.CsumStart)
if cSumAt+1 >= len(pkt) { if cSumAt+1 >= len(pkt) {
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt)) return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
} }
@@ -151,18 +144,13 @@ func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
} }
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a // SegmentTCP walks a TSO superpacket pkt, yielding each segment as a
// slice into pkt itself. Per-segment plaintext is laid out by stamping a // slice into pkt itself. Per-segment plaintext is laid out by sliding a
// copy of the original L3+L4 header into pkt at offset i*gsoSize, where it // freshly-patched copy of the L3+L4 header into pkt at offset i*gsoSize,
// sits immediately before that segment's payload chunk in the original // where it sits immediately before that segment's payload chunk in the
// buffer. The stamp is destructive but harmless: iter i's header write lands // original buffer. The slide is destructive: iter i's header write overwrites
// on pkt[i*G : i*G+hdrLen], which is the tail of seg_{i-1}'s payload (already // the last hdrLen bytes of seg_{i-1}'s payload, which is dead by the time
// consumed) and ends exactly where seg_i's payload begins, so it never clobbers // the next iteration begins. pkt is consumed by this call and must not be
// live payload — this holds even when gsoSize < hdrLen. The header bytes are // inspected by the caller after the final yield.
// sourced from a pristine snapshot taken before the loop (savedHdr), NOT from
// pkt[:hdrLen], because when gsoSize < hdrLen the stamps would otherwise
// overwrite the leading header in place and every stamp after the first would
// copy corrupted bytes. pkt is consumed by this call and must not be inspected
// by the caller after the final yield.
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error { func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
if gsoSizeU == 0 { if gsoSizeU == 0 {
return fmt.Errorf("gso_size is zero") return fmt.Errorf("gso_size is zero")
@@ -173,9 +161,6 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
headerLen := int(hdrLenU) headerLen := int(hdrLenU)
csumStart := int(csumStartU) csumStart := int(csumStartU)
if headerLen > maxSegHdrLen {
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
}
isV4 := pkt[0]>>4 == 4 isV4 := pkt[0]>>4 == 4
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4 tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
@@ -220,13 +205,6 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0)) baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
} }
// Snapshot the pristine L3+L4 header once. Every segment's header is
// stamped from this copy, so overlapping stamps (gsoSize < headerLen)
// can never corrupt the source. The variable fields (seq/flags/cksum/
// totalLen/id) captured here are stale but are overwritten per segment.
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := 0; i < numSeg; i++ { for i := 0; i < numSeg; i++ {
segStart := i * gsoSize segStart := i * gsoSize
segEnd := segStart + gsoSize segEnd := segStart + gsoSize
@@ -237,13 +215,14 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
segLen := headerLen + segPayLen segLen := headerLen + segPayLen
headerOff := i * gsoSize headerOff := i * gsoSize
// Stamp the header into place immediately before this segment's // Slide the header into place immediately before this segment's
// payload, sourced from the pristine snapshot. Iter 0's header is // payload. Iter 0's header is already at pkt[:headerLen]; for
// already at pkt[:headerLen] (identical to savedHdr), so only i ≥ 1 // i ≥ 1 we copy from there. The constant-byte fields of pkt[:headerLen]
// needs the stamp. The per-segment patches below overwrite the // survive iter 0's in-place patches (only seq/flags/cksum/totalLen/id
// variable fields. // are touched), and iter 0's stale variable-field values are
// overwritten by the per-segment patches below.
if i > 0 { if i > 0 {
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen]) copy(pkt[headerOff:headerOff+headerLen], pkt[:headerLen])
} }
seg := pkt[headerOff : headerOff+segLen] seg := pkt[headerOff : headerOff+segLen]
@@ -290,13 +269,11 @@ func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
return nil return nil
} }
// SegmentUDP walks a USO superpacket, stamping a per-segment-patched copy of // SegmentUDP walks a USO superpacket, sliding a per-segment-patched
// the original L3+L4 header into pkt at offset i*gsoSize and yielding // L3+L4 header into pkt at offset i*gsoSize and yielding pkt[i*G:i*G+segLen]
// pkt[i*G:i*G+segLen] to the caller. Per-segment patches are total_len + // to the caller. Per-segment patches are total_len + IPv4 csum (or IPv6
// IPv4 csum (or IPv6 payload_len) plus the UDP length and checksum. pkt is // payload_len) plus the UDP length and checksum. pkt is consumed
// consumed destructively; see SegmentTCP for the layout reasoning, including // destructively; see SegmentTCP for the layout reasoning.
// why the header is stamped from a pristine snapshot rather than pkt[:hdrLen]
// (correctness when gsoSize < hdrLen).
// //
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not // UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
// bump it), which is why the IP-level per-segment work is limited to // bump it), which is why the IP-level per-segment work is limited to
@@ -312,9 +289,6 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
isV4 := pkt[0]>>4 == 4 isV4 := pkt[0]>>4 == 4
headerLen := int(hdrLenU) headerLen := int(hdrLenU)
csumStart := int(csumStartU) csumStart := int(csumStartU)
if headerLen > maxSegHdrLen {
return fmt.Errorf("header len %d exceeds max %d", headerLen, maxSegHdrLen)
}
if headerLen-csumStart != udpHeaderLen { if headerLen-csumStart != udpHeaderLen {
return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart) return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart)
} }
@@ -353,12 +327,6 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0)) baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
} }
// Snapshot the pristine L3+L4 header once and stamp every segment from
// it; see SegmentTCP for why sourcing from pkt[:headerLen] corrupts
// segments when gsoSize < headerLen.
var savedHdr [maxSegHdrLen]byte
copy(savedHdr[:headerLen], pkt[:headerLen])
for i := 0; i < numSeg; i++ { for i := 0; i < numSeg; i++ {
segStart := i * gsoSize segStart := i * gsoSize
segEnd := segStart + gsoSize segEnd := segStart + gsoSize
@@ -370,7 +338,7 @@ func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg
headerOff := i * gsoSize headerOff := i * gsoSize
if i > 0 { if i > 0 {
copy(pkt[headerOff:headerOff+headerLen], savedHdr[:headerLen]) copy(pkt[headerOff:headerOff+headerLen], pkt[:headerLen])
} }
seg := pkt[headerOff : headerOff+segLen] seg := pkt[headerOff : headerOff+segLen]
-335
View File
@@ -1,335 +0,0 @@
//go:build linux && !android
// +build linux,!android
package virtio
import (
"bytes"
"encoding/binary"
"testing"
"golang.org/x/sys/unix"
"github.com/slackhq/nebula/overlay/checksum"
)
// verifyChecksum confirms that the one's-complement sum across b, seeded with
// a folded pseudo-header sum, equals all-ones (a valid on-wire checksum).
// A corrupted header stamped into a segment makes this fail even when the
// checksum field itself was computed from the (pristine) base sums, because
// the bytes the receiver would sum no longer match what was checksummed.
func verifyChecksum(b []byte, pseudo uint16) bool {
return checksum.Checksum(b, pseudo) == 0xffff
}
// pseudoHeaderIPv4 folds the TCP/UDP pseudo-header sum from a segment's own
// address and length fields, used to independently verify its L4 checksum.
func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 {
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
s += uint32(proto) + uint32(l4Len)
s = (s & 0xffff) + (s >> 16)
s = (s & 0xffff) + (s >> 16)
return uint16(s)
}
// buildTCPv4Super constructs a synthetic IPv4/TCP TSO superpacket with a
// payload of payLen bytes and returns it alongside the header fields the
// segmenter needs. The header is a fixed 40 bytes (20 IPv4 + 20 TCP).
func buildTCPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
const ipLen = 20
const tcpLen = 20
pkt = make([]byte, ipLen+tcpLen+payLen)
// IPv4 header.
pkt[0] = 0x45 // version 4, IHL 5
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+tcpLen+payLen))
binary.BigEndian.PutUint16(pkt[4:6], 0x4242) // ID
pkt[8] = 64 // TTL
pkt[9] = unix.IPPROTO_TCP
copy(pkt[12:16], []byte{10, 0, 0, 1}) // src
copy(pkt[16:20], []byte{10, 0, 0, 2}) // dst
// TCP header.
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 80) // dport
binary.BigEndian.PutUint32(pkt[24:28], 10000) // seq
binary.BigEndian.PutUint32(pkt[28:32], 20000) // ack
pkt[32] = 0x50 // data offset 5 words
pkt[33] = 0x18 // ACK | PSH
binary.BigEndian.PutUint16(pkt[34:36], 65535) // window
for i := 0; i < payLen; i++ {
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + tcpLen, ipLen
}
// buildUDPv4Super constructs a synthetic IPv4/UDP USO superpacket with a
// payload of payLen bytes. Header is a fixed 28 bytes (20 IPv4 + 8 UDP).
func buildUDPv4Super(payLen int) (pkt []byte, hdrLen, csumStart uint16) {
const ipLen = 20
const udpLen = 8
pkt = make([]byte, ipLen+udpLen+payLen)
pkt[0] = 0x45
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen))
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
pkt[8] = 64
pkt[9] = unix.IPPROTO_UDP
copy(pkt[12:16], []byte{10, 0, 0, 1})
copy(pkt[16:20], []byte{10, 0, 0, 2})
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
for i := 0; i < payLen; i++ {
pkt[ipLen+udpLen+i] = byte(i & 0xff)
}
return pkt, ipLen + udpLen, ipLen
}
// collectTCP segments a fresh copy of pkt and returns each segment as an
// independent slice so assertions can run after segmentation completes.
func collectTCP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
t.Helper()
work := append([]byte(nil), pkt...)
var out [][]byte
err := SegmentTCP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
out = append(out, append([]byte(nil), seg...))
return nil
})
if err != nil {
t.Fatalf("SegmentTCP: %v", err)
}
return out
}
func collectUDP(t *testing.T, pkt []byte, hdrLen, csumStart, gsoSize uint16) [][]byte {
t.Helper()
work := append([]byte(nil), pkt...)
var out [][]byte
err := SegmentUDP(work, hdrLen, csumStart, gsoSize, func(seg []byte) error {
out = append(out, append([]byte(nil), seg...))
return nil
})
if err != nil {
t.Fatalf("SegmentUDP: %v", err)
}
return out
}
// TestSegmentTCPHeaderNotCorrupted is the regression test for the in-place
// header-slide bug: when gsoSize < headerLen the old code stamped each
// segment's header from pkt[:headerLen], which had already been overwritten
// by the previous segment's overlapping stamp, so segments 2..n carried a
// corrupted header (garbage src/dst/ports/seq). Every segment must instead
// carry the ORIGINAL constant header fields with correct per-segment seq.
func TestSegmentTCPHeaderNotCorrupted(t *testing.T) {
const origSeq = 10000
cases := []struct {
name string
payLen int
gsoSize uint16
}{
// gsoSize (8) < headerLen (40): the bug's trigger. Even split.
{"small-gso-even", 40, 8},
// gsoSize (8) < headerLen (40) with a short final segment.
{"small-gso-odd-tail", 44, 8},
// gsoSize (100) >= headerLen (40): the normal path, must still work.
{"normal-gso", 250, 100},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
pkt, hdrLen, csumStart := buildTCPv4Super(tc.payLen)
gso := int(tc.gsoSize)
wantSeg := (tc.payLen + gso - 1) / gso
segs := collectTCP(t, pkt, hdrLen, csumStart, tc.gsoSize)
if len(segs) != wantSeg {
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
}
off := 0
for i, seg := range segs {
// Constant header fields must be identical to the original in
// EVERY segment. These are exactly the bytes the old code
// corrupted in segments 2..n.
if got := seg[0]; got != 0x45 {
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
}
if seg[9] != unix.IPPROTO_TCP {
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_TCP)
}
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
}
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
}
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
t.Errorf("seg %d: sport=%d want 12345", i, sport)
}
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 80 {
t.Errorf("seg %d: dport=%d want 80", i, dport)
}
if ack := binary.BigEndian.Uint32(seg[28:32]); ack != 20000 {
t.Errorf("seg %d: ack=%d want 20000", i, ack)
}
if seg[32] != 0x50 {
t.Errorf("seg %d: data-offset byte=%#x want 0x50", i, seg[32])
}
// Per-segment seq must advance by the payload offset.
segStart := i * gso
if seq := binary.BigEndian.Uint32(seg[24:28]); seq != uint32(origSeq+segStart) {
t.Errorf("seg %d: seq=%d want %d", i, seq, origSeq+segStart)
}
// Payload bytes must be the original contiguous slice.
segPayLen := len(seg) - int(hdrLen)
wantPay := make([]byte, segPayLen)
for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
t.Errorf("seg %d: payload mismatch", i)
}
off += segPayLen
// End-to-end: the stamped header must checksum-verify. A
// corrupted header fails here because the written checksum was
// derived from the pristine header.
if !verifyChecksum(seg[:20], 0) {
t.Errorf("seg %d: bad IPv4 header checksum", i)
}
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, len(seg)-20)
if !verifyChecksum(seg[20:], psum) {
t.Errorf("seg %d: bad TCP checksum", i)
}
}
})
}
}
// TestCorrectHdrLenChecksumBound guards the checksum-field bounds check in
// CorrectHdrLen. The checksum field sits at CsumStart+CsumOffset, so the check
// must be computed from CsumStart+CsumOffset — NOT CsumStart+CsumStart, a
// regression that doubled CsumStart and thus over-tightened the bound (since
// CsumOffset, 6 for UDP / 16 for TCP, is always < CsumStart >= 20). That bogus
// bound spuriously rejected valid small USO superpackets in decodeRead.
func TestCorrectHdrLenChecksumBound(t *testing.T) {
// A valid IPv4 USO superpacket: 20B IPv4 + 8B UDP + two 6-byte segments
// (payload 12) = 40 bytes total. CsumStart=20, CsumOffset=6, so the UDP
// checksum field lives at bytes 26..27, comfortably inside the 40-byte
// packet. The OLD formula computed cSumAt = CsumStart+CsumStart = 40 and
// rejected on cSumAt+1 (41) >= len(pkt) (40); the fix (CsumStart+CsumOffset
// = 26) accepts. This case FAILS against the CsumStart+CsumStart regression.
t.Run("valid-small-uso-accepted", func(t *testing.T) {
pkt, _, csumStart := buildUDPv4Super(12) // total len 40
hdr := Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
GSOSize: 6, // two 6-byte segments
CsumStart: csumStart,
CsumOffset: 6,
}
if err := CorrectHdrLen(pkt, &hdr); err != nil {
t.Fatalf("CorrectHdrLen rejected a valid 40-byte USO superpacket: %v", err)
}
if hdr.HdrLen != csumStart+udpHeaderLen {
t.Errorf("HdrLen = %d, want %d", hdr.HdrLen, csumStart+udpHeaderLen)
}
})
// A genuinely-too-short packet: CsumStart=20, CsumOffset=6 means the
// checksum field would end at byte 27, but the packet is only 25 bytes
// (CsumStart+CsumOffset+2 = 28 > 25). CorrectHdrLen must still reject it.
t.Run("too-short-rejected", func(t *testing.T) {
pkt := make([]byte, 25)
pkt[0] = 0x45 // IPv4, IHL 5
hdr := Hdr{
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
GSOSize: 6,
CsumStart: 20,
CsumOffset: 6,
}
if err := CorrectHdrLen(pkt, &hdr); err == nil {
t.Fatalf("CorrectHdrLen accepted a too-short (25-byte) packet")
}
})
}
// TestSegmentUDPHeaderNotCorrupted is the USO counterpart: SegmentUDP performs
// the same header stamp and must be correct when gsoSize < headerLen.
func TestSegmentUDPHeaderNotCorrupted(t *testing.T) {
cases := []struct {
name string
payLen int
gsoSize uint16
}{
{"small-gso-even", 40, 8},
{"small-gso-odd-tail", 44, 8},
{"normal-gso", 250, 100},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
pkt, hdrLen, csumStart := buildUDPv4Super(tc.payLen)
gso := int(tc.gsoSize)
wantSeg := (tc.payLen + gso - 1) / gso
segs := collectUDP(t, pkt, hdrLen, csumStart, tc.gsoSize)
if len(segs) != wantSeg {
t.Fatalf("got %d segments, want %d", len(segs), wantSeg)
}
off := 0
for i, seg := range segs {
if got := seg[0]; got != 0x45 {
t.Errorf("seg %d: version/IHL byte=%#x want 0x45", i, got)
}
if seg[9] != unix.IPPROTO_UDP {
t.Errorf("seg %d: proto=%d want %d", i, seg[9], unix.IPPROTO_UDP)
}
if !bytes.Equal(seg[12:16], []byte{10, 0, 0, 1}) {
t.Errorf("seg %d: src=%v want [10 0 0 1]", i, seg[12:16])
}
if !bytes.Equal(seg[16:20], []byte{10, 0, 0, 2}) {
t.Errorf("seg %d: dst=%v want [10 0 0 2]", i, seg[16:20])
}
if sport := binary.BigEndian.Uint16(seg[20:22]); sport != 12345 {
t.Errorf("seg %d: sport=%d want 12345", i, sport)
}
if dport := binary.BigEndian.Uint16(seg[22:24]); dport != 53 {
t.Errorf("seg %d: dport=%d want 53", i, dport)
}
// UDP-GSO keeps the same IPv4 ID across every segment.
if id := binary.BigEndian.Uint16(seg[4:6]); id != 0x4242 {
t.Errorf("seg %d: ip id=%#x want 0x4242", i, id)
}
segPayLen := len(seg) - int(hdrLen)
if udpLen := binary.BigEndian.Uint16(seg[24:26]); udpLen != uint16(8+segPayLen) {
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+segPayLen)
}
wantPay := make([]byte, segPayLen)
for k := 0; k < segPayLen; k++ {
wantPay[k] = byte((off + k) & 0xff)
}
if !bytes.Equal(seg[hdrLen:], wantPay) {
t.Errorf("seg %d: payload mismatch", i)
}
off += segPayLen
if !verifyChecksum(seg[:20], 0) {
t.Errorf("seg %d: bad IPv4 header checksum", i)
}
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, len(seg)-20)
if !verifyChecksum(seg[20:], psum) {
t.Errorf("seg %d: bad UDP checksum", i)
}
}
})
}
}
+33 -5
View File
@@ -19,12 +19,32 @@ import (
) )
type tun struct { type tun struct {
io.ReadWriteCloser rwc io.ReadWriteCloser
fd int fd int
vpnNetworks []netip.Prefix vpnNetworks []netip.Prefix
Routes atomic.Pointer[[]Route] Routes atomic.Pointer[[]Route]
routeTree atomic.Pointer[bart.Table[routing.Gateways]] routeTree atomic.Pointer[bart.Table[routing.Gateways]]
l *slog.Logger l *slog.Logger
readBuf []byte
batchRet [1]tio.Packet
}
func (t *tun) Read() ([]tio.Packet, error) {
n, err := t.rwc.Read(t.readBuf)
if err != nil {
return nil, err
}
t.batchRet[0] = tio.Packet{Bytes: t.readBuf[:n]}
return t.batchRet[:], nil
}
func (t *tun) Write(p []byte) (int, error) {
return t.rwc.Write(p)
}
func (t *tun) Close() error {
return t.rwc.Close()
} }
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) { func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
@@ -33,15 +53,15 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun") file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
t := &tun{ t := &tun{
ReadWriteCloser: file, rwc: file,
fd: deviceFd, fd: deviceFd,
vpnNetworks: vpnNetworks, vpnNetworks: vpnNetworks,
l: l, l: l,
readBuf: make([]byte, defaultBatchBufSize),
} }
err := t.reload(c, true) err := t.reload(c, true)
if err != nil { if err != nil {
_ = file.Close()
return nil, err return nil, err
} }
@@ -97,6 +117,14 @@ func (t *tun) Name() string {
return "android" return "android"
} }
func (t *tun) Queues(int) ([]tio.Queue, error) { func (t *tun) SupportsMultiqueue() bool {
return []tio.Queue{tio.NewSingleQueue(t, defaultBatchBufSize)}, nil return false
}
func (t *tun) NewMultiQueueReader() error {
return fmt.Errorf("TODO: multiqueue not implemented for android")
}
func (t *tun) Readers() []tio.Queue {
return []tio.Queue{t}
} }

Some files were not shown because too many files have changed in this diff Show More