mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 09:16:57 +02:00
Compare commits
16 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f5ddff5ca1 | |||
| 400cbc26a1 | |||
| 01b31360df | |||
| 5bdf645b0b | |||
| 0375aff451 | |||
| 6cb00c613c | |||
| 40b4ae7fb4 | |||
| cf51b6dfd7 | |||
| fe93ebd017 | |||
| 961ddbfbc1 | |||
| 67bd9e848a | |||
| bc3f5d0400 | |||
| aef8e39cc4 | |||
| 69863d6c81 | |||
| 5d35351437 | |||
| f95857b4c3 |
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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}" .
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -268,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
|
||||||
@@ -280,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 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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
+7
-32
@@ -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,16 +261,14 @@ 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
|
||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,13 +57,11 @@ 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 != "" {
|
||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+20
-42
@@ -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,10 +266,16 @@ 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 == "" {
|
||||||
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
*sf.outKeyPath = *sf.name + ".key"
|
||||||
return fmt.Errorf("refusing to overwrite existing cert: %s", *sf.outCertPath)
|
}
|
||||||
}
|
|
||||||
|
if *sf.outCertPath == "" {
|
||||||
|
*sf.outCertPath = *sf.name + ".crt"
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := os.Stat(*sf.outCertPath); err == nil {
|
||||||
|
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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
assert.Empty(t, 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())
|
|
||||||
|
|
||||||
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())
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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"))
|
|
||||||
}
|
|
||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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`)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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,12 +61,9 @@ 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)
|
||||||
@@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
+6
-8
@@ -50,12 +50,9 @@ 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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -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
@@ -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,
|
||||||
|
|||||||
+1
-1
@@ -10,7 +10,7 @@ import (
|
|||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 8192
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey noiseutil.CipherState
|
||||||
|
|||||||
+25
-51
@@ -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 {
|
||||||
|
|||||||
@@ -1,296 +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/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},
|
|
||||||
routines: 1,
|
|
||||||
hostMap: newHostMap(l),
|
|
||||||
lightHouse: lh,
|
|
||||||
l: l,
|
|
||||||
}
|
|
||||||
f.wg.Add(1)
|
|
||||||
|
|
||||||
return &Control{
|
|
||||||
state: StateReady,
|
|
||||||
f: f,
|
|
||||||
l: l,
|
|
||||||
ctx: ctx,
|
|
||||||
cancel: cancel,
|
|
||||||
}, dev, conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_StopBeforeStart(t *testing.T) {
|
|
||||||
c, dev, conn := newReadyControl(t)
|
|
||||||
|
|
||||||
// A Stop on a never started control must release everything Main acquired
|
|
||||||
c.Stop()
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
|
||||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
|
||||||
require.ErrorIs(t, c.ctx.Err(), context.Canceled, "the service context should have been cancelled")
|
|
||||||
|
|
||||||
// Wait must return promptly now that the resources are released
|
|
||||||
require.NoError(t, c.Wait())
|
|
||||||
|
|
||||||
// A stopped control can never be started
|
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
|
||||||
|
|
||||||
// A second Stop is a harmless no-op
|
|
||||||
c.Stop()
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
require.NoError(t, c.Wait())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_WaitBlocksUntilStop(t *testing.T) {
|
|
||||||
c, _, _ := newReadyControl(t)
|
|
||||||
|
|
||||||
done := make(chan error, 1)
|
|
||||||
go func() { done <- c.Wait() }()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
t.Fatal("Wait returned before Stop")
|
|
||||||
case <-time.After(50 * time.Millisecond):
|
|
||||||
}
|
|
||||||
|
|
||||||
c.Stop()
|
|
||||||
select {
|
|
||||||
case err := <-done:
|
|
||||||
require.NoError(t, err)
|
|
||||||
case <-time.After(time.Second):
|
|
||||||
t.Fatal("Wait did not return after Stop")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type fakeConn struct {
|
|
||||||
closed bool
|
|
||||||
rebinds int
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *fakeConn) Rebind() error { c.rebinds++; return nil }
|
|
||||||
func (c *fakeConn) LocalAddr() (netip.AddrPort, error) { return netip.AddrPort{}, nil }
|
|
||||||
func (c *fakeConn) ListenOut(_ udp.EncReader) error { return nil }
|
|
||||||
func (c *fakeConn) WriteTo(_ []byte, _ netip.AddrPort) error { return nil }
|
|
||||||
func (c *fakeConn) ReloadConfig(_ *config.C) {}
|
|
||||||
func (c *fakeConn) SupportsMultipleReaders() bool { return true }
|
|
||||||
func (c *fakeConn) Close() error { c.closed = true; return nil }
|
|
||||||
|
|
||||||
type multiqueueDevice struct {
|
|
||||||
*fakeDevice
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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},
|
|
||||||
routines: 2,
|
|
||||||
l: test.NewLogger(),
|
|
||||||
}
|
|
||||||
f.wg.Add(1)
|
|
||||||
|
|
||||||
c := &Control{
|
|
||||||
state: StateReady,
|
|
||||||
f: f,
|
|
||||||
l: test.NewLogger(),
|
|
||||||
ctx: ctx,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
|
|
||||||
// The second reader fails to open, everything must be released
|
|
||||||
require.Error(t, c.Start())
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
|
||||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
|
||||||
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
|
||||||
|
|
||||||
// And Wait must not hang on the construction token
|
|
||||||
require.NoError(t, c.Wait())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInterface_CloseIsIdempotent(t *testing.T) {
|
|
||||||
dev := newFakeDevice()
|
|
||||||
f := &Interface{
|
|
||||||
inside: dev,
|
|
||||||
l: test.NewLogger(),
|
|
||||||
}
|
|
||||||
f.wg.Add(1)
|
|
||||||
|
|
||||||
require.NoError(t, f.Close())
|
|
||||||
assert.True(t, dev.closed)
|
|
||||||
|
|
||||||
// A second Close must not double release the wg token or the device
|
|
||||||
require.NoError(t, f.Close())
|
|
||||||
require.NoError(t, f.wait())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_FatalErrorReportsThroughWait(t *testing.T) {
|
|
||||||
c, dev, conn := newReadyControl(t)
|
|
||||||
|
|
||||||
// Mirror what Start wires up, without needing real packet readers
|
|
||||||
c.f.triggerShutdown = c.Stop
|
|
||||||
c.state = StateStarted
|
|
||||||
|
|
||||||
boom := errors.New("boom")
|
|
||||||
c.f.onFatal(boom)
|
|
||||||
|
|
||||||
require.ErrorIs(t, c.Wait(), boom)
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
assert.True(t, dev.closed)
|
|
||||||
assert.True(t, conn.closed)
|
|
||||||
|
|
||||||
// A second fatal error must not fire the shutdown again or replace the first
|
|
||||||
c.f.onFatal(errors.New("later"))
|
|
||||||
require.ErrorIs(t, c.Wait(), boom)
|
|
||||||
|
|
||||||
// Wait stays factual, a Stop after the death does not mask the error
|
|
||||||
c.Stop()
|
|
||||||
require.ErrorIs(t, c.Wait(), boom)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_ConcurrentStopAndStart(t *testing.T) {
|
|
||||||
c, _, _ := newReadyControl(t)
|
|
||||||
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for i := 0; i < 2; i++ {
|
|
||||||
wg.Go(func() { c.Stop() })
|
|
||||||
}
|
|
||||||
wg.Go(func() { _ = c.Start() })
|
|
||||||
wg.Go(func() {
|
|
||||||
_ = c.Wait()
|
|
||||||
// A returned Wait must always observe the final state, no matter how
|
|
||||||
// the race resolved
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
})
|
|
||||||
wg.Wait()
|
|
||||||
|
|
||||||
// However the race resolves, the control must end fully stopped with no
|
|
||||||
// panic and Wait must observe the final state
|
|
||||||
require.NoError(t, c.Wait())
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_StartStopLifecycle(t *testing.T) {
|
|
||||||
c, dev, conn := newReadyControl(t)
|
|
||||||
|
|
||||||
require.NoError(t, c.Start())
|
|
||||||
assert.Equal(t, StateStarted, c.State())
|
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStarted)
|
|
||||||
|
|
||||||
// Stop must unpark the reader blocked in the device and release everything
|
|
||||||
c.Stop()
|
|
||||||
assert.Equal(t, StateStopped, c.State())
|
|
||||||
assert.True(t, dev.closed, "the tun device should have been closed")
|
|
||||||
assert.True(t, conn.closed, "the udp socket should have been closed")
|
|
||||||
require.ErrorIs(t, c.ctx.Err(), context.Canceled)
|
|
||||||
|
|
||||||
// The reader drained off a closed device, that is not a fatal error
|
|
||||||
require.NoError(t, c.Wait())
|
|
||||||
require.ErrorIs(t, c.Start(), ErrAlreadyStopped)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestControl_RebindIsGatedByState(t *testing.T) {
|
|
||||||
c, _, conn := newReadyControl(t)
|
|
||||||
|
|
||||||
// A rebind before Start reaches nothing, the interface is not up
|
|
||||||
c.RebindUDPServer()
|
|
||||||
assert.Equal(t, 0, conn.rebinds, "rebind before start must be a no-op")
|
|
||||||
|
|
||||||
require.NoError(t, c.Start())
|
|
||||||
c.RebindUDPServer()
|
|
||||||
assert.Equal(t, 1, conn.rebinds, "rebind while started must reach the conn")
|
|
||||||
|
|
||||||
// A rebind racing a completed stop must not touch the closed conn
|
|
||||||
c.Stop()
|
|
||||||
require.NoError(t, c.Wait())
|
|
||||||
c.RebindUDPServer()
|
|
||||||
assert.Equal(t, 1, conn.rebinds, "rebind after stop must be a no-op")
|
|
||||||
}
|
|
||||||
+6
-161
@@ -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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+1
-1
@@ -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))
|
||||||
|
|||||||
+26
-92
@@ -11,21 +11,19 @@ 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"
|
||||||
)
|
)
|
||||||
|
|
||||||
type dnsServer struct {
|
type dnsServer struct {
|
||||||
sync.RWMutex
|
sync.RWMutex
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
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) {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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()
|
|
||||||
}
|
|
||||||
|
|||||||
+6
-101
@@ -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")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,125 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestInnerECN(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
pkt []byte
|
||||||
|
want byte
|
||||||
|
}{
|
||||||
|
{"empty", nil, 0},
|
||||||
|
{"v4_NotECT", v4WithToS(0x00), 0x00},
|
||||||
|
{"v4_ECT0", v4WithToS(0x02), 0x02},
|
||||||
|
{"v4_ECT1", v4WithToS(0x01), 0x01},
|
||||||
|
{"v4_CE", v4WithToS(0x03), 0x03},
|
||||||
|
{"v4_DSCP_then_NotECT", v4WithToS(0x88 | 0x00), 0x00},
|
||||||
|
{"v4_DSCP_then_CE", v4WithToS(0x88 | 0x03), 0x03},
|
||||||
|
{"v6_NotECT", v6WithTC(0x00), 0x00},
|
||||||
|
{"v6_ECT0", v6WithTC(0x02), 0x02},
|
||||||
|
{"v6_CE", v6WithTC(0x03), 0x03},
|
||||||
|
{"v6_DSCP_then_CE", v6WithTC(0x88 | 0x03), 0x03},
|
||||||
|
{"unknown_version", []byte{0xa5, 0xff}, 0},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
got := innerECN(c.pkt)
|
||||||
|
if got != c.want {
|
||||||
|
t.Errorf("innerECN=0x%02x want 0x%02x", got, c.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// v4WithToS returns a 2-byte slice tall enough for innerECN: byte 0 carries
|
||||||
|
// version=4 in the high nibble, byte 1 is the full ToS so we exercise both
|
||||||
|
// the DSCP and ECN portions through the byte 1 mask.
|
||||||
|
func v4WithToS(tos byte) []byte {
|
||||||
|
return []byte{0x45, tos}
|
||||||
|
}
|
||||||
|
|
||||||
|
// v6WithTC builds a 2-byte slice that places a known traffic class value
|
||||||
|
// across bytes 0 (high nibble of TC) and 1 (low nibble of TC). innerECN
|
||||||
|
// extracts ECN as (b[1]>>4)&0x03, which corresponds to TC[1:0].
|
||||||
|
func v6WithTC(tc byte) []byte {
|
||||||
|
return []byte{0x60 | (tc>>4)&0x0f, (tc & 0x0f) << 4}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyOuterECN(t *testing.T) {
|
||||||
|
silent := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||||
|
hi := &HostInfo{}
|
||||||
|
|
||||||
|
// Build a v4 packet helper with a given inner ECN field.
|
||||||
|
v4 := func(innerECN byte) []byte {
|
||||||
|
// 20-byte minimal IPv4 header with ToS = innerECN (DSCP zeroed).
|
||||||
|
return []byte{
|
||||||
|
0x45, innerECN, 0, 28,
|
||||||
|
0, 0, 0x40, 0,
|
||||||
|
64, 6, 0, 0,
|
||||||
|
10, 0, 0, 1,
|
||||||
|
10, 0, 0, 2,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Build a v6 packet helper with a given inner ECN field. ECN occupies
|
||||||
|
// TC[1:0] which sit at byte 1 mask 0x30.
|
||||||
|
v6 := func(innerECN byte) []byte {
|
||||||
|
// 40-byte minimal IPv6 header with TC[1:0] = innerECN.
|
||||||
|
pkt := make([]byte, 40)
|
||||||
|
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
||||||
|
pkt[1] = (innerECN & 0x03) << 4 // TC[3:0]: low 2 bits = ECN, top 2 = DSCP-low (0)
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
type cell struct {
|
||||||
|
outer byte
|
||||||
|
inner byte
|
||||||
|
wantECN byte
|
||||||
|
wantSame bool // expect inner unchanged (true => verify the byte didn't move)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RFC 6040 normal-mode combine table. Only outer==CE causes mutation.
|
||||||
|
table := []cell{
|
||||||
|
{ecnNotECT, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnNotECT, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnNotECT, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnNotECT, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnECT0, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnECT0, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnECT0, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnECT0, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnECT1, ecnNotECT, ecnNotECT, true},
|
||||||
|
{ecnECT1, ecnECT0, ecnECT0, true},
|
||||||
|
{ecnECT1, ecnECT1, ecnECT1, true},
|
||||||
|
{ecnECT1, ecnCE, ecnCE, true},
|
||||||
|
|
||||||
|
{ecnCE, ecnNotECT, ecnNotECT, true}, // legacy: log, leave alone
|
||||||
|
{ecnCE, ecnECT0, ecnCE, false}, // CE folded in
|
||||||
|
{ecnCE, ecnECT1, ecnCE, false},
|
||||||
|
{ecnCE, ecnCE, ecnCE, true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, c := range table {
|
||||||
|
t.Run("v4", func(t *testing.T) {
|
||||||
|
pkt := v4(c.inner)
|
||||||
|
applyOuterECN(pkt, c.outer, hi, silent)
|
||||||
|
got := pkt[1] & 0x03
|
||||||
|
if got != c.wantECN {
|
||||||
|
t.Errorf("v4 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
t.Run("v6", func(t *testing.T) {
|
||||||
|
pkt := v6(c.inner)
|
||||||
|
applyOuterECN(pkt, c.outer, hi, silent)
|
||||||
|
got := (pkt[1] >> 4) & 0x03
|
||||||
|
if got != c.wantECN {
|
||||||
|
t.Errorf("v6 outer=0x%02x inner=0x%02x: got 0x%02x want 0x%02x", c.outer, c.inner, got, c.wantECN)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+1
-15
@@ -254,20 +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
|
|
||||||
# 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.
|
|
||||||
#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. 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
|
||||||
@@ -411,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
|
||||||
|
|
||||||
|
|||||||
+81
-79
@@ -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
|
||||||
@@ -58,9 +58,8 @@ type Firewall struct {
|
|||||||
routableNetworks *bart.Lite
|
routableNetworks *bart.Lite
|
||||||
|
|
||||||
// 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,25 +158,26 @@ 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(),
|
||||||
TCPTimeout: tcpTimeout,
|
TCPTimeout: tcpTimeout,
|
||||||
UDPTimeout: UDPTimeout,
|
UDPTimeout: UDPTimeout,
|
||||||
DefaultTimeout: defaultTimeout,
|
DefaultTimeout: defaultTimeout,
|
||||||
routableNetworks: routableNetworks,
|
routableNetworks: routableNetworks,
|
||||||
assignedNetworks: assignedNetworks,
|
assignedNetworks: assignedNetworks,
|
||||||
unsafeNetworks: unsafeNetworks,
|
hasUnsafeNetworks: hasUnsafeNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
|
|
||||||
incomingMetrics: firewallMetrics{
|
incomingMetrics: firewallMetrics{
|
||||||
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
droppedLocalAddr: metrics.GetOrRegisterCounter("firewall.incoming.dropped.local_addr", nil),
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
|
|||||||
+8
-4
@@ -5,11 +5,15 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
)
|
)
|
||||||
|
|
||||||
// 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
|
||||||
@@ -56,8 +60,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
|||||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||||
c.l.Debug("resetting conntrack cache", "len", ll)
|
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/logging"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
@@ -22,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
|
||||||
@@ -30,27 +31,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
c := newFixedTicker(t, l, 3)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
c := newFixedTicker(t, l, 2)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
c := newFixedTicker(t, l, 5)
|
||||||
c.Get()
|
c.Get()
|
||||||
@@ -60,7 +61,7 @@ func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
c := newFixedTicker(t, l, 0)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|||||||
@@ -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
@@ -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()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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=
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+12
-4
@@ -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) {
|
||||||
@@ -967,6 +974,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
for _, cp := range hh.packetStore {
|
for _, cp := range hh.packetStore {
|
||||||
|
//todo use a sendbatcher
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
}
|
}
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
|||||||
+118
-171
@@ -56,20 +56,11 @@ type Relay struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type HostMap struct {
|
type HostMap struct {
|
||||||
sync.RWMutex //Because we concurrently read and write to our maps
|
sync.RWMutex //Because we concurrently read and write to our maps
|
||||||
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,66 +400,85 @@ 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
|
|
||||||
|
// Unlink this hostinfo
|
||||||
|
if hostinfo.prev != nil {
|
||||||
|
hostinfo.prev.next = hostinfo.next
|
||||||
|
}
|
||||||
|
if hostinfo.next != nil {
|
||||||
|
hostinfo.next.prev = hostinfo.prev
|
||||||
|
}
|
||||||
|
|
||||||
|
// If there wasn't a previous primary then clear out any links
|
||||||
|
if oldHostinfo == nil {
|
||||||
|
hostinfo.next = nil
|
||||||
|
hostinfo.prev = nil
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Relink the hostinfo as primary
|
||||||
|
hostinfo.next = oldHostinfo
|
||||||
|
oldHostinfo.prev = hostinfo
|
||||||
|
hostinfo.prev = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// unlockedDeleteHostInfo removes hostinfo from every one of its address lists and from the index
|
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) {
|
||||||
// maps. It returns true if this was the last hostinfo for all of its addresses (we no longer have
|
|
||||||
// any tunnel to the peer), which the caller uses to decide whether to clear learned lighthouse
|
|
||||||
// state and disestablish relays.
|
|
||||||
func (hm *HostMap) unlockedDeleteHostInfo(hostinfo *HostInfo) bool {
|
|
||||||
// Remove this hostinfo from each of its address lists. The lists are independent, so a
|
|
||||||
// sibling is never promoted to an address it does not own and no other list is touched.
|
|
||||||
final := true
|
|
||||||
for _, addr := range hostinfo.vpnAddrs {
|
for _, addr := range hostinfo.vpnAddrs {
|
||||||
if list, ok := hm.moreHosts[addr]; ok {
|
h := hm.Hosts[addr]
|
||||||
list = removeHostInfo(list, hostinfo)
|
for h != nil {
|
||||||
hm.unlockedSetHostsForAddr(addr, list)
|
if h == hostinfo {
|
||||||
if len(list) > 0 {
|
hm.unlockedInnerDeleteHostInfo(h, addr)
|
||||||
final = false
|
|
||||||
}
|
|
||||||
} else if existing, ok := hm.Hosts[addr]; ok {
|
|
||||||
if existing == hostinfo {
|
|
||||||
// Common case, the only hostinfo for this address. moreHosts has no entry to clean up.
|
|
||||||
delete(hm.Hosts, addr)
|
|
||||||
} else {
|
|
||||||
// We don't hold this address but another hostinfo does, we still have a tunnel to the peer
|
|
||||||
final = false
|
|
||||||
}
|
}
|
||||||
|
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)
|
||||||
|
if len(hm.Hosts) == 0 {
|
||||||
|
hm.Hosts = map[netip.Addr]*HostInfo{}
|
||||||
|
}
|
||||||
|
|
||||||
|
if hostinfo.next != nil {
|
||||||
|
// We had more than 1 hostinfo at this vpn addr, promote the next in the list to primary
|
||||||
|
hm.Hosts[addr] = hostinfo.next
|
||||||
|
// It is primary, there is no previous hostinfo now
|
||||||
|
hostinfo.next.prev = nil
|
||||||
|
}
|
||||||
|
|
||||||
|
} else {
|
||||||
|
// Relink if we were in the middle of multiple hostinfos for this vpn addr
|
||||||
|
if hostinfo.prev != nil {
|
||||||
|
hostinfo.prev.next = hostinfo.next
|
||||||
|
}
|
||||||
|
|
||||||
|
if hostinfo.next != nil {
|
||||||
|
hostinfo.next.prev = hostinfo.prev
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Go maps never shrink their buckets, replace fully drained maps so a node that churned
|
hostinfo.next = nil
|
||||||
// through a large peer count gives the memory back. Same idiom as the index maps below.
|
hostinfo.prev = nil
|
||||||
if len(hm.Hosts) == 0 {
|
|
||||||
hm.Hosts = map[netip.Addr]*HostInfo{}
|
|
||||||
}
|
|
||||||
if len(hm.moreHosts) == 0 {
|
|
||||||
hm.moreHosts = map[netip.Addr][]*HostInfo{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// The remote index uses index ids outside our control so lets make sure we are only removing
|
// The remote index 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
|
||||||
@@ -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 _, targetIp := range targetIps {
|
for h != nil {
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
for _, targetIp := range targetIps {
|
||||||
if ok && r.State == Established {
|
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
||||||
return h, r, nil
|
if ok && r.State == Established {
|
||||||
}
|
return h, r, nil
|
||||||
}
|
|
||||||
|
|
||||||
if list, ok := hm.moreHosts[relayHostIp]; ok {
|
|
||||||
// list[0] is the primary we already checked
|
|
||||||
for _, h := range list[1:] {
|
|
||||||
for _, targetIp := range targetIps {
|
|
||||||
r, ok := h.relayState.QueryRelayForByIp(targetIp)
|
|
||||||
if ok && r.State == Established {
|
|
||||||
return h, r, nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
h = h.next
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, nil, errors.New("unable to find host with relay")
|
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 {
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
for h != nil {
|
||||||
|
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 {
|
||||||
h.relayState.UpdateRelayForByIpState(hi.vpnAddrs[0], Disestablished)
|
for h != nil {
|
||||||
|
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 {
|
hm.Hosts[vpnAddr] = hostinfo
|
||||||
// Common case, the first hostinfo for this address. moreHosts stays empty.
|
|
||||||
hm.Hosts[vpnAddr] = hostinfo
|
if existing != nil && existing != hostinfo {
|
||||||
return
|
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
@@ -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) {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -9,12 +10,26 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
err := newPacket(packet, false, fwPacket)
|
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||||
if err != nil {
|
// only valid until the next Read on that queue. Every consumer below
|
||||||
|
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
||||||
|
// synchronously; do not retain pkt outside this call. If a future
|
||||||
|
// caller needs to keep the packet, use pkt.Clone() to detach it from
|
||||||
|
// the borrow.
|
||||||
|
//
|
||||||
|
// 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
|
||||||
|
// the same 5-tuple every segment will share, so a single parse +
|
||||||
|
// firewall check covers the whole superpacket.
|
||||||
|
packet := pkt.Bytes
|
||||||
|
var parsed batch.RxParsed
|
||||||
|
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,
|
||||||
@@ -24,6 +39,8 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
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) {
|
||||||
@@ -37,7 +54,14 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
_, err := f.queues[q].Write(packet)
|
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
|
||||||
|
// A self-forwarded superpacket would be re-handed to the
|
||||||
|
// kernel as one giant blob; segment first so the loopback
|
||||||
|
// path sees one IP datagram per Write.
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
_, werr := f.readers[q].Write(seg)
|
||||||
|
return werr
|
||||||
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -53,11 +77,23 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||||
|
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||||
|
// so retaining segments past the loop is safe.
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Failed to segment superpacket for handshake cache",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -71,12 +107,11 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
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.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -86,8 +121,151 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
|
||||||
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
|
f.connectionManager.Out(hostinfo)
|
||||||
|
|
||||||
|
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Unlock()
|
||||||
|
}
|
||||||
|
if encErr != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||||
|
"error", encErr,
|
||||||
|
"udpAddr", hostinfo.remote,
|
||||||
|
"counter", c,
|
||||||
|
)
|
||||||
|
// Skip this segment; the rest of the superpacket can still
|
||||||
|
// go out — TCP will retransmit anything we drop here.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendInsideMessage encrypts a firewall-approved inside packet (or every
|
||||||
|
// segment of a TSO/USO superpacket) into the caller's batch slot for
|
||||||
|
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||||
|
// kernel-supplied superpacket bytes never get written into a separate
|
||||||
|
// scratch arena: SegmentSuperpacket builds each segment's plaintext in
|
||||||
|
// segScratch[:segLen] in turn, and we encrypt directly into a fresh
|
||||||
|
// SendBatch slot.
|
||||||
|
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ecnEnabled := f.ecnEnabled.Load()
|
||||||
|
if hostinfo.lastRebindCount != f.rebindCount {
|
||||||
|
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||||
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
|
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||||
|
hostinfo.lastRebindCount = f.rebindCount
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hostinfo.remote.IsValid() { //the relay path
|
||||||
|
//first, find our relay hostinfo:
|
||||||
|
var relayHostInfo *HostInfo
|
||||||
|
var relay *Relay
|
||||||
|
var err error
|
||||||
|
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
||||||
|
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.relayState.DeleteRelay(relayIP)
|
||||||
|
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
||||||
|
"relay", relayIP,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if relayHostInfo == nil || relay == nil {
|
||||||
|
//failure already logged
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err = tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
//relay header + header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305) + relay tag
|
||||||
|
scratch := sendBatch.Reserve(header.Len + header.Len + len(seg) + 16 + 16)
|
||||||
|
|
||||||
|
innerPacket := f.sendInsideEncrypt(hostinfo, ci, seg, scratch[header.Len:], nb)
|
||||||
|
if innerPacket == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
//now we need to do a relay-encrypt:
|
||||||
|
toSend, err := f.prepareSendVia(relayHostInfo, relay, innerPacket, nb, scratch, true)
|
||||||
|
if err != nil {
|
||||||
|
//already logged
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecnEnabled {
|
||||||
|
ecn = innerECN(seg)
|
||||||
|
}
|
||||||
|
sendBatch.Commit(toSend, relayHostInfo.remote, ecn)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
||||||
|
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
||||||
|
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
||||||
|
|
||||||
|
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
||||||
|
if out == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var ecn byte
|
||||||
|
if ecnEnabled {
|
||||||
|
ecn = innerECN(seg)
|
||||||
|
}
|
||||||
|
sendBatch.Commit(out, hostinfo.remote, ecn)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
||||||
|
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
||||||
|
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
||||||
|
func innerECN(pkt []byte) byte {
|
||||||
|
if len(pkt) < 2 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
switch pkt[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
return pkt[1] & 0x03
|
||||||
|
case 6:
|
||||||
|
return (pkt[1] >> 4) & 0x03
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -216,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",
|
||||||
@@ -275,21 +454,13 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
|||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
|
||||||
// handshake messages to peers through relay tunnels.
|
|
||||||
// via is the HostInfo through which the message is relayed.
|
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
|
||||||
// q indicates which writer to use to send the packet.
|
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
|
||||||
relay *Relay,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) {
|
) ([]byte, error) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
@@ -311,7 +482,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return
|
return nil, io.ErrShortBuffer
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -331,20 +502,39 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
return
|
return nil, err
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.GetRemote())
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||||
|
// to the payload for the ultimate target host, making this a useful method for sending
|
||||||
|
// handshake messages to peers through relay tunnels.
|
||||||
|
// via is the HostInfo through which the message is relayed.
|
||||||
|
// ad is the plaintext data to authenticate, but not encrypt
|
||||||
|
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||||
|
// out is a buffer used to store the result of the Encrypt operation
|
||||||
|
// q indicates which writer to use to send the packet.
|
||||||
|
func (f *Interface) SendVia(via *HostInfo,
|
||||||
|
relay *Relay,
|
||||||
|
ad,
|
||||||
|
nb,
|
||||||
|
out []byte,
|
||||||
|
nocopy bool,
|
||||||
|
) {
|
||||||
|
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||||
|
err = f.writers[0].WriteTo(toSend, via.remote)
|
||||||
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)
|
||||||
}
|
}
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||||
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 {
|
||||||
@@ -391,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
|
||||||
}
|
}
|
||||||
@@ -403,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,
|
||||||
|
|||||||
+131
-131
@@ -7,22 +7,21 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime"
|
"runtime"
|
||||||
"slices"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
|
"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"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/slackhq/nebula/util"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const mtu = 9001
|
const mtu = 9001
|
||||||
@@ -55,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 UDP sends on one XPS-selected NIC TX ring so per-flow
|
|
||||||
// packets stay ordered on the wire.
|
|
||||||
PinThreads bool
|
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
@@ -89,13 +83,13 @@ 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
|
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
||||||
// a CPU at all (tun.pin_threads, default true). When false, threads are
|
// inside.go copies the inner ECN onto the outer carrier on encap and
|
||||||
// left free to migrate as on stock nebula.
|
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
||||||
pinThreads bool
|
// via tunnels.ecn (default true).
|
||||||
|
ecnEnabled atomic.Bool
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
|
|
||||||
tryPromoteEvery atomic.Uint32
|
tryPromoteEvery atomic.Uint32
|
||||||
@@ -113,8 +107,12 @@ type Interface struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
queues []tio.Queue
|
readers []tio.Queue
|
||||||
wg sync.WaitGroup
|
// batchers is one per tun queue, wrapping readers[i].
|
||||||
|
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||||
|
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||||
|
batchers []batch.RxBatcher
|
||||||
|
wg sync.WaitGroup
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
// nil means "no fatal error" (yet)
|
// nil means "no fatal error" (yet)
|
||||||
@@ -212,6 +210,8 @@ 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),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -221,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,
|
||||||
@@ -239,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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -264,37 +260,48 @@ func (f *Interface) activate() error {
|
|||||||
"boringcrypto", boringEnabled(),
|
"boringcrypto", boringEnabled(),
|
||||||
)
|
)
|
||||||
|
|
||||||
if f.routines > 1 && !f.outside.SupportsMultipleReaders() {
|
if f.routines > 1 {
|
||||||
f.routines = 1
|
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
|
||||||
f.l.Warn("multiple udp readers are not supported on this platform, falling back to a single routine")
|
f.routines = 1
|
||||||
|
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))
|
||||||
|
|
||||||
// On error the caller owns the cleanup, Control.Start cancels the service context
|
// Prepare n tun queues
|
||||||
// before releasing our resources so a waiter never observes a live context
|
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 {
|
||||||
|
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
||||||
|
// is on, everything else (and either lane disabled) falls
|
||||||
|
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
||||||
|
// reaches the TUN.
|
||||||
|
f.batchers[i] = batch.NewMultiCoalescer(f.readers[i], caps.TSO, caps.USO)
|
||||||
|
} else {
|
||||||
|
f.batchers[i] = batch.NewPassthrough(f.readers[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
f.wg.Add(1) // for us to wait on Close() to return
|
||||||
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() {
|
||||||
@@ -305,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 {
|
||||||
|
f.wg.Wait()
|
||||||
func (f *Interface) wait() error {
|
if e := f.fatalErr.Load(); e != nil {
|
||||||
f.wg.Wait()
|
return *e
|
||||||
if e := f.fatalErr.Load(); e != nil {
|
}
|
||||||
return *e
|
return nil
|
||||||
}
|
}, nil
|
||||||
return 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
|
||||||
@@ -340,19 +346,25 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
plaintext := make([]byte, udp.MTU)
|
|
||||||
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)
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
plaintext := f.batchers[i].Reserve(len(payload))
|
||||||
})
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, parsedRx, lhh, nb, i, ctCache.Get(), meta)
|
||||||
|
}
|
||||||
|
|
||||||
// An error after teardown began is shutdown noise, the closed flag covers resources
|
flusher := func() {
|
||||||
// Close releases itself and the cancelled ctx covers ones torn down by their owners
|
if err := f.batchers[i].Flush(); err != nil {
|
||||||
// reacting to it, like the user device pipes
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
if err != nil && !f.closed.Load() && f.ctx.Err() == nil {
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err := li.ListenOut(listener, flusher)
|
||||||
|
|
||||||
|
if err != nil && !f.closed.Load() {
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -360,40 +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 UDP send from this goroutine going through
|
// Pin this goroutine to one CPU. LockOSThread alone keeps the goroutine
|
||||||
// the same TX ring on the nic (XPS selects the ring by CPU), so the wire sees per-flow order. Skip entirely
|
// on a single OS thread but the kernel can still migrate that thread
|
||||||
// when tun.pin_threads is false.
|
// across CPUs — XPS reads smp_processor_id() at sendmmsg time and picks
|
||||||
if f.pinThreads {
|
// the TX ring from the current CPU's xps_cpus map, so an unpinned
|
||||||
var cpu int
|
// thread bouncing between CPUs spreads one nebula flow's packets across
|
||||||
if n := len(f.cpuAffinity); n > 0 {
|
// multiple TX rings, which the rings then drain at independent rates
|
||||||
// Explicit tun.cpu_affinity list wins; parseCpuAffinity already
|
// and the wire delivers reordered.
|
||||||
// validated the entries against the allowed CPU set.
|
//
|
||||||
cpu = f.cpuAffinity[i%n]
|
// Pinning keeps every sendmmsg from this goroutine going through the
|
||||||
} else if allowed, err := util.AllowedCPUs(); err == nil && len(allowed) > 0 {
|
// same TX ring, so the wire sees per-flow order. Cost: less scheduler
|
||||||
// Default: spread queues across the CPUs we're actually allowed to
|
// flexibility — if i % NumCPU collides between two TUN reader
|
||||||
// run on. Under a cpuset/taskset mask these aren't 0..NumCPU-1, so
|
// goroutines they share a CPU.
|
||||||
// i % NumCPU would pick unrunnable IDs and every pin would fail.
|
cpu := i % runtime.NumCPU()
|
||||||
cpu = allowed[i%len(allowed)]
|
if n := len(f.cpuAffinity); n > 0 {
|
||||||
} else {
|
cpu = f.cpuAffinity[i%n]
|
||||||
cpu = i % runtime.NumCPU()
|
|
||||||
}
|
|
||||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
|
||||||
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||||
out := make([]byte, mtu)
|
f.l.Warn("failed to pin tun reader to CPU", "queue", i, "cpu", cpu, "err", err)
|
||||||
|
}
|
||||||
|
rejectBuf := make([]byte, mtu)
|
||||||
|
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, udp.MTU+32)
|
||||||
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)
|
||||||
}
|
}
|
||||||
@@ -401,9 +410,10 @@ func (f *Interface) listenIn(queue tio.Queue, i int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for _, pkt := range pkts {
|
for _, pkt := range pkts {
|
||||||
// borrowed: pkt.Bytes is owned by the queue and only valid until
|
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||||
// the next Read; consumeInsidePacket reads it synchronously.
|
}
|
||||||
f.consumeInsidePacket(pkt.Bytes, fwPacket, nb, out, i, conntrackCache.Get())
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -416,6 +426,7 @@ func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
|||||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||||
c.RegisterReloadCallback(f.reloadMisc)
|
c.RegisterReloadCallback(f.reloadMisc)
|
||||||
|
c.RegisterReloadCallback(f.reloadEcn)
|
||||||
|
|
||||||
for _, udpConn := range f.writers {
|
for _, udpConn := range f.writers {
|
||||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||||
@@ -433,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
|
||||||
@@ -548,6 +550,20 @@ func (f *Interface) reloadMisc(c *config.C) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
||||||
|
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
||||||
|
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
||||||
|
func (f *Interface) reloadEcn(c *config.C) {
|
||||||
|
initial := c.InitialLoad()
|
||||||
|
if initial || c.HasChanged("tunnels.ecn") {
|
||||||
|
v := c.GetBool("tunnels.ecn", true)
|
||||||
|
f.ecnEnabled.Store(v)
|
||||||
|
if !initial {
|
||||||
|
f.l.Info("tunnels.ecn changed", "enabled", v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||||
ticker := time.NewTicker(i)
|
ticker := time.NewTicker(i)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
@@ -558,34 +574,26 @@ 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() {
|
|
||||||
f.firewall.EmitStats()
|
|
||||||
f.handshakeManager.EmitStats()
|
|
||||||
udpStats()
|
|
||||||
|
|
||||||
certState := f.pki.getCertState()
|
|
||||||
defaultCrt := certState.GetDefaultCertificate()
|
|
||||||
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
|
||||||
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
|
||||||
|
|
||||||
// Report the max certificate version we are capable of using
|
|
||||||
if certState.v2Cert != nil {
|
|
||||||
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
|
||||||
} else {
|
|
||||||
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 {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
emit()
|
f.firewall.EmitStats()
|
||||||
|
f.handshakeManager.EmitStats()
|
||||||
|
udpStats()
|
||||||
|
|
||||||
|
certState := f.pki.getCertState()
|
||||||
|
defaultCrt := certState.GetDefaultCertificate()
|
||||||
|
certExpirationGauge.Update(int64(defaultCrt.NotAfter().Sub(time.Now()) / time.Second))
|
||||||
|
certInitiatingVersion.Update(int64(defaultCrt.Version()))
|
||||||
|
|
||||||
|
// Report the max certificate version we are capable of using
|
||||||
|
if certState.v2Cert != nil {
|
||||||
|
certMaxVersion.Update(int64(certState.v2Cert.Version()))
|
||||||
|
} else {
|
||||||
|
certMaxVersion.Update(int64(certState.v1Cert.Version()))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -598,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 {
|
||||||
@@ -622,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...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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())
|
|
||||||
}
|
|
||||||
@@ -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")
|
|
||||||
}
|
|
||||||
+19
-290
@@ -4,54 +4,26 @@ 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 packet[9] {
|
||||||
switch version {
|
case 6: // tcp
|
||||||
case ipv4.Version:
|
return ipv4CreateRejectTCPPacket(packet, out)
|
||||||
if len(packet) < ipv4.HeaderLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Do not send reject packets for non-first fragments
|
|
||||||
if packet[6]&0x1f != 0 || packet[7] != 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
switch packet[9] {
|
|
||||||
case 6: // tcp
|
|
||||||
return ipv4CreateRejectTCPPacket(packet, out)
|
|
||||||
default:
|
|
||||||
return ipv4CreateRejectICMPPacket(packet, out)
|
|
||||||
}
|
|
||||||
case ipv6.Version:
|
|
||||||
if len(packet) < ipv6.HeaderLen {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return ipv6CreateRejectPacket(packet, out)
|
|
||||||
default:
|
default:
|
||||||
return nil
|
return ipv4CreateRejectICMPPacket(packet, out)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -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) {
|
||||||
@@ -104,14 +71,14 @@ 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
|
||||||
icmpOut[5] = 0 // .
|
icmpOut[5] = 0 // .
|
||||||
icmpOut[6] = 0 // .
|
icmpOut[6] = 0 // .
|
||||||
icmpOut[7] = 0 // .
|
icmpOut[7] = 0 // .
|
||||||
|
|
||||||
// Copy original IP header and first 8 bytes as body
|
// Copy original IP header and first 8 bytes as body
|
||||||
copy(icmpOut[8:], packet[:packetLen])
|
copy(icmpOut[8:], packet[:packetLen])
|
||||||
@@ -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
-404
@@ -1,13 +1,11 @@
|
|||||||
package iputil
|
package iputil
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"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) {
|
||||||
@@ -45,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)
|
||||||
@@ -73,404 +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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
|
|||||||
+6
-31
@@ -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,9 +1477,7 @@ 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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,6 +36,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
buildVersion = moduleVersion()
|
buildVersion = moduleVersion()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
//todo no merge
|
||||||
|
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 {
|
||||||
b, err := yaml.Marshal(c.Settings)
|
b, err := yaml.Marshal(c.Settings)
|
||||||
@@ -131,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
|
||||||
@@ -206,7 +200,7 @@ 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)
|
||||||
}
|
}
|
||||||
@@ -233,7 +227,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
punchy: punchy,
|
punchy: punchy,
|
||||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||||
CpuAffinity: parseCpuAffinity(c, l, routines),
|
CpuAffinity: parseCpuAffinity(c, l, routines),
|
||||||
PinThreads: c.GetBool("tun.pin_threads", true),
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -251,6 +244,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
ifce.reloadDisconnectInvalid(c)
|
ifce.reloadDisconnectInvalid(c)
|
||||||
ifce.reloadSendRecvError(c)
|
ifce.reloadSendRecvError(c)
|
||||||
ifce.reloadAcceptRecvError(c)
|
ifce.reloadAcceptRecvError(c)
|
||||||
|
ifce.reloadEcn(c)
|
||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
@@ -287,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 {
|
||||||
@@ -307,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
|
||||||
@@ -330,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)
|
||||||
|
|||||||
@@ -1,51 +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 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))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+91
-219
@@ -2,27 +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"
|
||||||
"golang.org/x/net/ipv4"
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
)
|
"github.com/slackhq/nebula/udp"
|
||||||
|
|
||||||
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) {
|
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
|
||||||
@@ -110,7 +103,7 @@ 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)
|
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)
|
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) {
|
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)
|
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
|
||||||
}
|
}
|
||||||
@@ -236,7 +223,7 @@ func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender,
|
|||||||
switch targetRelay.Type {
|
switch targetRelay.Type {
|
||||||
case ForwardingType:
|
case ForwardingType:
|
||||||
// Forward this packet through the relay tunnel
|
// Forward this packet through the relay tunnel
|
||||||
// Find the target HostInfo
|
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||||
case TerminalType:
|
case TerminalType:
|
||||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -518,8 +331,68 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []
|
|||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
||||||
err := newPacket(out, true, fwPacket)
|
const (
|
||||||
|
ecnNotECT = 0x00
|
||||||
|
ecnECT1 = 0x01
|
||||||
|
ecnECT0 = 0x02
|
||||||
|
ecnCE = 0x03
|
||||||
|
)
|
||||||
|
|
||||||
|
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
||||||
|
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
||||||
|
// codepoints are advisory only and leave the inner unchanged.
|
||||||
|
//
|
||||||
|
// Merge cases (outer × inner → action):
|
||||||
|
//
|
||||||
|
// outer != CE : no-op (inner is authoritative)
|
||||||
|
// outer == CE, inner Not-ECT : log; cannot propagate to a non-ECN host
|
||||||
|
// outer == CE, inner ECT/CE : rewrite inner ECN to CE
|
||||||
|
func applyOuterECN(pkt []byte, outerECN byte, hostinfo *HostInfo, l *slog.Logger) {
|
||||||
|
if outerECN&ecnCE != ecnCE || len(pkt) < 2 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
switch pkt[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
switch pkt[1] & 0x03 {
|
||||||
|
case ecnNotECT:
|
||||||
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||||
|
}
|
||||||
|
case ecnCE:
|
||||||
|
// Already CE.
|
||||||
|
default:
|
||||||
|
pkt[1] = (pkt[1] &^ 0x03) | ecnCE
|
||||||
|
}
|
||||||
|
case 6:
|
||||||
|
switch (pkt[1] >> 4) & 0x03 {
|
||||||
|
case ecnNotECT:
|
||||||
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(l).Debug("RFC 6040: outer CE on inner Not-ECT, leaving inner unchanged")
|
||||||
|
}
|
||||||
|
case ecnCE:
|
||||||
|
// Already CE.
|
||||||
|
default:
|
||||||
|
pkt[1] = (pkt[1] &^ 0x30) | (ecnCE << 4)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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
|
||||||
|
// underlay into the inner header before firewall + TUN write. Other
|
||||||
|
// outer codepoints are advisory only — we keep the inner unchanged.
|
||||||
|
if f.ecnEnabled.Load() {
|
||||||
|
applyOuterECN(out, meta.OuterECN, hostinfo, f.l)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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,
|
||||||
@@ -528,7 +401,7 @@ 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
|
// This gives us a buffer to build the reject packet in
|
||||||
@@ -542,7 +415,7 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err = f.queues[q].Write(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)
|
||||||
}
|
}
|
||||||
@@ -589,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
@@ -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")
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,35 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import "net/netip"
|
||||||
|
|
||||||
|
type RxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// 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
|
||||||
|
// CommitInbound is Commit with a hint produced by ParsePacket, so the
|
||||||
|
// batcher can skip the IP+L4 re-parse. Borrowed slice contract is the
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
type TxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt and records its destination plus the 2-bit
|
||||||
|
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
||||||
|
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
||||||
|
// to leave the outer ECN field unset.
|
||||||
|
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
||||||
|
// Flush emits every queued packet via the underlying batch writer in
|
||||||
|
// arrival order. Returns an errors.Join of one or more errors. After Flush returns,
|
||||||
|
// borrowed payload slices may be recycled.
|
||||||
|
Flush() error
|
||||||
|
}
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
|
)
|
||||||
|
|
||||||
|
// flowKey identifies a transport flow by {src, dst, sport, dport, family}.
|
||||||
|
// Comparable, so map lookups and linear scans over the slot list stay tight.
|
||||||
|
// Shared by the TCP and UDP coalescers; each coalescer keeps its own
|
||||||
|
// openSlots map, so a TCP and UDP flow on the same 5-tuple-without-proto
|
||||||
|
// never alias.
|
||||||
|
type flowKey struct {
|
||||||
|
src, dst [16]byte
|
||||||
|
sport, dport uint16
|
||||||
|
isV6 bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// initialSlots is the starting capacity of the slot pool. One flow per
|
||||||
|
// packet is the worst case so this matches a typical carrier-side
|
||||||
|
// recvmmsg batch on the encrypted UDP socket.
|
||||||
|
const initialSlots = 64
|
||||||
|
|
||||||
|
// parsedIP is the IP-level result of parseIPPrologue. The caller layers
|
||||||
|
// L4-specific parsing (TCP / UDP) on top.
|
||||||
|
type parsedIP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
// pkt is the original buffer trimmed to the IP-declared total length.
|
||||||
|
// Anything below the IP layer (transport parsers) should slice into
|
||||||
|
// pkt rather than the unbounded original.
|
||||||
|
pkt []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseIPPrologue extracts the IP-level fields the coalescers care about:
|
||||||
|
// IHL/payload length, version, src/dst addresses, and the L4 protocol byte.
|
||||||
|
// Returns ok=false for malformed input, IPv4 with options or fragmentation,
|
||||||
|
// or IPv6 with extension headers (all rejected by both coalescers in
|
||||||
|
// identical ways before this refactor).
|
||||||
|
//
|
||||||
|
// On success, p.pkt is len-trimmed to the IP-declared length so callers
|
||||||
|
// don't have to repeat the trim. wantProto is the IANA protocol number to
|
||||||
|
// require (6 for TCP, 17 for UDP); ok=false for any other value.
|
||||||
|
func parseIPPrologue(pkt []byte, wantProto byte) (parsedIP, bool) {
|
||||||
|
var p parsedIP
|
||||||
|
if len(pkt) < 20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
v := pkt[0] >> 4
|
||||||
|
switch v {
|
||||||
|
case 4:
|
||||||
|
ihl := int(pkt[0]&0x0f) * 4
|
||||||
|
if ihl != 20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[9] != wantProto {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
// Reject actual fragmentation (MF or non-zero frag offset).
|
||||||
|
if binary.BigEndian.Uint16(pkt[6:8])&0x3fff != 0 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
totalLen := int(binary.BigEndian.Uint16(pkt[2:4]))
|
||||||
|
if totalLen > len(pkt) || totalLen < ihl {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.ipHdrLen = 20
|
||||||
|
p.fk.isV6 = false
|
||||||
|
copy(p.fk.src[:4], pkt[12:16])
|
||||||
|
copy(p.fk.dst[:4], pkt[16:20])
|
||||||
|
p.pkt = pkt[:totalLen]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[6] != wantProto {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
payloadLen := int(binary.BigEndian.Uint16(pkt[4:6]))
|
||||||
|
if 40+payloadLen > len(pkt) {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.ipHdrLen = 40
|
||||||
|
p.fk.isV6 = true
|
||||||
|
copy(p.fk.src[:], pkt[8:24])
|
||||||
|
copy(p.fk.dst[:], pkt[24:40])
|
||||||
|
p.pkt = pkt[:40+payloadLen]
|
||||||
|
default:
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// ipHeadersMatch compares the IP portion of two packet header prefixes for
|
||||||
|
// byte-for-byte equality on every field that must be identical across
|
||||||
|
// coalesced segments. Size/IPID/IPCsum and the 2-bit IP-level ECN field are
|
||||||
|
// masked out — the appendPayload step merges CE into the seed.
|
||||||
|
//
|
||||||
|
// The transport (L4) portion of the header is checked separately by the
|
||||||
|
// per-protocol matcher.
|
||||||
|
func ipHeadersMatch(a, b []byte, isV6 bool) bool {
|
||||||
|
if isV6 {
|
||||||
|
// 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.
|
||||||
|
// ECN lives in TC[1:0] = byte 1 mask 0x30. Skip [4:6] payload_len.
|
||||||
|
if a[0] != b[0] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if a[1]&^0x30 != b[1]&^0x30 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[2:4], b[2:4]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:40], b[6:40]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// IPv4: byte 0 = version/IHL, byte 1 = DSCP(6)|ECN(2),
|
||||||
|
// [6:10] flags/fragoff/TTL/proto, [12:20] src+dst.
|
||||||
|
// Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||||
|
if a[0] != b[0] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if a[1]&^0x03 != b[1]&^0x03 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:10], b[6:10]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[12:20], b[12:20]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// mergeECNIntoSeed ORs the 2-bit IP-level ECN field of pkt's IP header
|
||||||
|
// onto the seed's IP header, so a CE mark on any coalesced segment
|
||||||
|
// propagates to the final superpacket. (CE is 0b11; ORing yields CE if
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// reserveFromBacking implements the Reserve half of the RxBatcher contract
|
||||||
|
// shared by TCP and UDP coalescers. The backing slice grows on demand;
|
||||||
|
// already-committed slices reference the old array and remain valid until
|
||||||
|
// Flush resets backing.
|
||||||
|
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)
|
||||||
|
*backing = (*backing)[:start+sz]
|
||||||
|
return (*backing)[start : start+sz : start+sz]
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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))
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MultiCoalescer fans plaintext packets out to lane-specific batchers based
|
||||||
|
// on the IP/L4 protocol of the packet, sharing a single Reserve arena
|
||||||
|
// across lanes so the caller's allocation pattern is unchanged.
|
||||||
|
//
|
||||||
|
// Lanes are processed independently: the TCP coalescer only sees TCP, the
|
||||||
|
// UDP coalescer only sees UDP, and the passthrough lane handles everything
|
||||||
|
// else. Per-flow arrival order is preserved because a single 5-tuple only
|
||||||
|
// ever lands in one lane and each lane preserves its own slot order.
|
||||||
|
//
|
||||||
|
// Cross-lane order is NOT preserved across the TCP/UDP/passthrough split.
|
||||||
|
// This is acceptable because the carrier-side recvmmsg path already
|
||||||
|
// stable-sorts by (peer, message counter) before delivering plaintext
|
||||||
|
// here, so replay-window invariants are unaffected, and apps observe
|
||||||
|
// correct per-flow ordering — which is all the IP layer guarantees anyway.
|
||||||
|
// Do not "fix" this by interleaving lane outputs at flush time; that
|
||||||
|
// negates the entire point of coalescing (each lane needs to see runs of
|
||||||
|
// adjacent same-flow packets to coalesce them).
|
||||||
|
type MultiCoalescer struct {
|
||||||
|
tcp *TCPCoalescer
|
||||||
|
udp *UDPCoalescer
|
||||||
|
pt *Passthrough
|
||||||
|
|
||||||
|
// arena shared across all lanes so a single Reserve grows one backing
|
||||||
|
// slice; lane Commit calls borrow into this same arena.
|
||||||
|
backing []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
|
// likewise gates UDP coalescing (only enable when USO was negotiated).
|
||||||
|
// Either lane disabled redirects its traffic into the passthrough lane.
|
||||||
|
func NewMultiCoalescer(w io.Writer, tcpEnabled, udpEnabled bool) *MultiCoalescer {
|
||||||
|
m := &MultiCoalescer{
|
||||||
|
pt: NewPassthrough(w),
|
||||||
|
backing: make([]byte, 0, initialSlots*65535),
|
||||||
|
}
|
||||||
|
if tcpEnabled {
|
||||||
|
m.tcp = NewTCPCoalescer(w)
|
||||||
|
}
|
||||||
|
if udpEnabled {
|
||||||
|
m.udp = NewUDPCoalescer(w)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *MultiCoalescer) Reserve(sz int) []byte {
|
||||||
|
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
|
||||||
|
// proto. Borrowed slice contract is identical to the single-lane batchers,
|
||||||
|
// pkt must remain valid until the next Flush.
|
||||||
|
//
|
||||||
|
// On the success path the IP/TCP-or-UDP parse happens here once and the
|
||||||
|
// parsed struct is handed to the lane via commitParsed so the lane doesn't
|
||||||
|
// re-walk the header.
|
||||||
|
func (m *MultiCoalescer) Commit(pkt []byte) error {
|
||||||
|
if len(pkt) < 20 {
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
v := pkt[0] >> 4
|
||||||
|
var proto byte
|
||||||
|
switch v {
|
||||||
|
case 4:
|
||||||
|
proto = pkt[9]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
proto = pkt[6]
|
||||||
|
default:
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
switch proto {
|
||||||
|
case ipProtoTCP:
|
||||||
|
if m.tcp != nil {
|
||||||
|
info, ok := parseTCPBase(pkt)
|
||||||
|
if !ok {
|
||||||
|
// Malformed/unsupported TCP shape (IP options, fragments, ...).
|
||||||
|
// Handle this via passthrough support in the TCP coalescer, to attempt to preserve flow order.
|
||||||
|
m.tcp.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.tcp.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
case ipProtoUDP:
|
||||||
|
if m.udp != nil {
|
||||||
|
info, ok := parseUDP(pkt)
|
||||||
|
if !ok {
|
||||||
|
m.udp.addPassthrough(pkt) //we could also m.pt.Commit() here I guess?
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return m.udp.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return m.pt.Commit(pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush drains every lane in a fixed order: TCP, UDP, passthrough. Errors
|
||||||
|
// 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 {
|
||||||
|
var errs []error
|
||||||
|
if m.tcp != nil {
|
||||||
|
if err := m.tcp.Flush(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if m.udp != nil {
|
||||||
|
if err := m.udp.Flush(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := m.pt.Flush(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
m.backing = m.backing[:0]
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -0,0 +1,94 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestMultiCoalescerRoutesByProto confirms TCP/UDP/other land in the right
|
||||||
|
// lane: TCP and UDP get coalesced when their lanes are enabled, anything
|
||||||
|
// else (ICMP here) falls through to plain Write.
|
||||||
|
func TestMultiCoalescerRoutesByProto(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, true, true)
|
||||||
|
|
||||||
|
tcpPay := make([]byte, 1200)
|
||||||
|
udpPay := make([]byte, 1200)
|
||||||
|
icmp := make([]byte, 28)
|
||||||
|
icmp[0] = 0x45
|
||||||
|
icmp[2] = 0
|
||||||
|
icmp[3] = 28
|
||||||
|
icmp[9] = 1
|
||||||
|
|
||||||
|
if err := m.Commit(buildTCPv4(1000, tcpAck, tcpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildTCPv4(2200, tcpAck, tcpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(2000, 53, udpPay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(icmp); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// 1 TCP super (2 segments) + 1 UDP super (2 segments) = 2 gso writes.
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 gso writes (one TCP + one UDP), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.writes) != 1 {
|
||||||
|
t.Fatalf("want 1 plain write (ICMP), got %d", len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMultiCoalescerDisabledUDPFallsThrough verifies that when the UDP lane
|
||||||
|
// is disabled (e.g. kernel doesn't support USO), UDP packets still reach
|
||||||
|
// the kernel via the passthrough lane rather than being lost.
|
||||||
|
func TestMultiCoalescerDisabledUDPFallsThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, true, false) // TSO on, USO off
|
||||||
|
|
||||||
|
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 0 {
|
||||||
|
t.Errorf("UDP must NOT be coalesced when USO disabled, got %d gso writes", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.writes) != 2 {
|
||||||
|
t.Errorf("UDP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestMultiCoalescerDisabledTCPFallsThrough mirrors the TSO=off case.
|
||||||
|
func TestMultiCoalescerDisabledTCPFallsThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
m := NewMultiCoalescer(w, false, true) // TSO off, USO on
|
||||||
|
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
if err := m.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := m.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 0 {
|
||||||
|
t.Errorf("TCP must NOT be coalesced when TSO disabled, got %d gso writes", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.writes) != 2 {
|
||||||
|
t.Errorf("TCP must pass through as 2 plain writes, got %d", len(w.writes))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||||
|
type Passthrough struct {
|
||||||
|
out io.Writer
|
||||||
|
slots [][]byte
|
||||||
|
backing []byte
|
||||||
|
cursor int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewPassthrough(w io.Writer) *Passthrough {
|
||||||
|
const baseNumSlots = 128
|
||||||
|
return &Passthrough{
|
||||||
|
out: w,
|
||||||
|
slots: make([][]byte, 0, baseNumSlots),
|
||||||
|
backing: make([]byte, 0, baseNumSlots*udp.MTU),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Reserve(sz int) []byte {
|
||||||
|
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 {
|
||||||
|
p.slots = append(p.slots, pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CommitInbound ignores the hint — Passthrough never coalesces, so there's
|
||||||
|
// no IP/L4 re-parse to skip. Present so Passthrough satisfies the RxBatcher
|
||||||
|
// interface alongside MultiCoalescer.
|
||||||
|
func (p *Passthrough) CommitInbound(pkt []byte, _ *RxParsed) error {
|
||||||
|
return p.Commit(pkt)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Flush() error {
|
||||||
|
var firstErr error
|
||||||
|
for _, s := range p.slots {
|
||||||
|
_, err := p.out.Write(s)
|
||||||
|
if err != nil && firstErr == nil {
|
||||||
|
firstErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
clear(p.slots)
|
||||||
|
p.slots = p.slots[:0]
|
||||||
|
p.backing = p.backing[:0]
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
@@ -0,0 +1,731 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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
|
||||||
|
|
||||||
|
// tcpCoalesceBufSize caps total bytes per superpacket. Mirrors the kernel's
|
||||||
|
// sk_gso_max_size of ~64KiB; anything beyond this would be rejected anyway.
|
||||||
|
const tcpCoalesceBufSize = 65535
|
||||||
|
|
||||||
|
// tcpCoalesceMaxSegs caps how many segments we'll coalesce into a single
|
||||||
|
// superpacket. Keeping this well below the kernel's TSO ceiling bounds
|
||||||
|
// latency.
|
||||||
|
const tcpCoalesceMaxSegs = 64
|
||||||
|
|
||||||
|
// tcpCoalesceHdrCap is the scratch space we copy a seed's IP+TCP header
|
||||||
|
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
||||||
|
const tcpCoalesceHdrCap = 100
|
||||||
|
|
||||||
|
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
||||||
|
// passthrough is true the slot holds a single borrowed packet that must be
|
||||||
|
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
||||||
|
// passthrough is false the slot is an in-progress coalesced superpacket:
|
||||||
|
// hdrBuf is a mutable copy of the seed's IP+TCP header (we patch total
|
||||||
|
// length and pseudo-header partial at flush), and payIovs are *borrowed*
|
||||||
|
// slices from the caller's plaintext buffers — no payload is ever copied.
|
||||||
|
// The caller (listenOut) must keep those buffers alive until Flush.
|
||||||
|
type coalesceSlot struct {
|
||||||
|
passthrough bool
|
||||||
|
rawPkt []byte // borrowed when passthrough
|
||||||
|
|
||||||
|
fk flowKey
|
||||||
|
hdrBuf [tcpCoalesceHdrCap]byte
|
||||||
|
hdrLen int
|
||||||
|
ipHdrLen int
|
||||||
|
isV6 bool
|
||||||
|
gsoSize int
|
||||||
|
numSeg int
|
||||||
|
totalPay int
|
||||||
|
nextSeq uint32
|
||||||
|
// psh closes the chain: set when the last-accepted segment had PSH or
|
||||||
|
// was sub-gsoSize. No further appends after that.
|
||||||
|
psh bool
|
||||||
|
payIovs [][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCPCoalescer accumulates adjacent in-flow TCP data segments across
|
||||||
|
// multiple concurrent flows and emits each flow's run as a single TSO
|
||||||
|
// superpacket via tio.GSOWriter. All output — coalesced or not — is
|
||||||
|
// deferred until Flush so arrival order is preserved on the wire. Owns
|
||||||
|
// no locks; one coalescer per TUN write queue.
|
||||||
|
type TCPCoalescer struct {
|
||||||
|
plainW io.Writer
|
||||||
|
gsoW tio.GSOWriter // nil when the queue doesn't support TSO
|
||||||
|
|
||||||
|
// slots is the ordered event queue. Flush walks it once and emits each
|
||||||
|
// entry as either a WriteGSO (coalesced) or a plainW.Write (passthrough).
|
||||||
|
slots []*coalesceSlot
|
||||||
|
// openSlots maps a flow key to its most recent non-sealed slot, so new
|
||||||
|
// segments can extend an in-progress superpacket in O(1). Slots are
|
||||||
|
// removed from this map when they close (PSH or short-last-segment),
|
||||||
|
// when a non-admissible packet for that flow arrives, or in Flush.
|
||||||
|
openSlots map[flowKey]*coalesceSlot
|
||||||
|
// lastSlot caches the most recently touched open slot. Steady-state
|
||||||
|
// bulk traffic is dominated by a single flow, so comparing the
|
||||||
|
// incoming key against the cached slot's own fk lets the hot path
|
||||||
|
// skip the map lookup (and the aeshash of a 38-byte key) entirely.
|
||||||
|
// Kept in lockstep with openSlots: nil whenever the slot it pointed
|
||||||
|
// at is removed/sealed.
|
||||||
|
lastSlot *coalesceSlot
|
||||||
|
pool []*coalesceSlot // free list for reuse
|
||||||
|
|
||||||
|
backing []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTCPCoalescer(w io.Writer) *TCPCoalescer {
|
||||||
|
c := &TCPCoalescer{
|
||||||
|
plainW: w,
|
||||||
|
slots: make([]*coalesceSlot, 0, initialSlots),
|
||||||
|
openSlots: make(map[flowKey]*coalesceSlot, initialSlots),
|
||||||
|
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||||
|
backing: make([]byte, 0, initialSlots*65535),
|
||||||
|
}
|
||||||
|
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
||||||
|
c.gsoW = gw
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsedTCP holds the fields extracted from a single parse so later steps
|
||||||
|
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||||
|
type parsedTCP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
tcpHdrLen int
|
||||||
|
hdrLen int
|
||||||
|
payLen int
|
||||||
|
seq uint32
|
||||||
|
flags byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseTCPBase extracts the flow key and IP/TCP offsets for any TCP packet,
|
||||||
|
// regardless of whether it's admissible for coalescing. Returns ok=false
|
||||||
|
// for non-TCP or malformed input. Accepts IPv4 (no options, no fragmentation)
|
||||||
|
// and IPv6 (no extension headers).
|
||||||
|
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
||||||
|
var p parsedTCP
|
||||||
|
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
||||||
|
if !ok {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
pkt = ip.pkt
|
||||||
|
p.fk = ip.fk
|
||||||
|
p.ipHdrLen = ip.ipHdrLen
|
||||||
|
|
||||||
|
if len(pkt) < p.ipHdrLen+20 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
tcpOff := int(pkt[p.ipHdrLen+12]>>4) * 4
|
||||||
|
if tcpOff < 20 || tcpOff > 60 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if len(pkt) < p.ipHdrLen+tcpOff {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.tcpHdrLen = tcpOff
|
||||||
|
p.hdrLen = p.ipHdrLen + tcpOff
|
||||||
|
p.payLen = len(pkt) - p.hdrLen
|
||||||
|
p.seq = binary.BigEndian.Uint32(pkt[p.ipHdrLen+4 : p.ipHdrLen+8])
|
||||||
|
p.flags = pkt[p.ipHdrLen+13]
|
||||||
|
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||||
|
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
// TCP flag bits (byte 13 of the TCP header). Only the bits actually consulted
|
||||||
|
// by the coalescer are named; FIN/SYN/RST/URG/CWR are rejected via the
|
||||||
|
// negative mask in coalesceable, not by name.
|
||||||
|
const (
|
||||||
|
tcpFlagPsh = 0x08
|
||||||
|
tcpFlagAck = 0x10
|
||||||
|
tcpFlagEce = 0x40
|
||||||
|
)
|
||||||
|
|
||||||
|
// coalesceable reports whether a parsed TCP segment is eligible for
|
||||||
|
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
||||||
|
// non-empty payload. CWR is excluded because it marks a one-shot
|
||||||
|
// congestion-window-reduced transition the receiver must observe at a
|
||||||
|
// segment boundary.
|
||||||
|
func (p parsedTCP) coalesceable() bool {
|
||||||
|
if p.flags&tcpFlagAck == 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return p.payLen > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
||||||
|
return reserveFromBacking(&c.backing, sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, ok := parseTCPBase(pkt)
|
||||||
|
if !ok {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
|
||||||
|
// commitParsed is the post-parse half of Commit. The caller must have
|
||||||
|
// already verified parseTCPBase succeeded (info is a valid TCP parse).
|
||||||
|
// Used by MultiCoalescer.Commit to avoid re-walking the IP/TCP header
|
||||||
|
// after the dispatcher has already done so.
|
||||||
|
func (c *TCPCoalescer) commitParsed(pkt []byte, info parsedTCP) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !info.coalesceable() {
|
||||||
|
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
||||||
|
// Seal this flow's open slot so later in-flow packets don't extend
|
||||||
|
// it and accidentally reorder past this passthrough.
|
||||||
|
if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Single-flow fast path: with only one open flow the cache hits every
|
||||||
|
// packet, and len(openSlots)==1 lets us skip the 38-byte fk compare
|
||||||
|
// when there are multiple flows in flight (where the hit rate would
|
||||||
|
// be ~0 and the compare is pure overhead).
|
||||||
|
var open *coalesceSlot
|
||||||
|
if last := c.lastSlot; last != nil && len(c.openSlots) == 1 && last.fk == info.fk {
|
||||||
|
open = last
|
||||||
|
} else {
|
||||||
|
open = c.openSlots[info.fk]
|
||||||
|
}
|
||||||
|
if open != nil {
|
||||||
|
if c.canAppend(open, pkt, info) {
|
||||||
|
c.appendPayload(open, pkt, info)
|
||||||
|
if open.psh {
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
c.lastSlot = nil
|
||||||
|
} else {
|
||||||
|
c.lastSlot = open
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
if c.lastSlot == open {
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
c.seed(pkt, info)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
c.reorderForFlush()
|
||||||
|
var first error
|
||||||
|
for _, s := range c.slots {
|
||||||
|
var err error
|
||||||
|
if s.passthrough {
|
||||||
|
_, err = c.plainW.Write(s.rawPkt)
|
||||||
|
} else {
|
||||||
|
err = c.flushSlot(s)
|
||||||
|
}
|
||||||
|
if err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
c.release(s)
|
||||||
|
}
|
||||||
|
clear(c.slots)
|
||||||
|
c.slots = c.slots[:0]
|
||||||
|
clear(c.openSlots)
|
||||||
|
c.lastSlot = nil
|
||||||
|
|
||||||
|
c.backing = c.backing[:0]
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) addPassthrough(pkt []byte) {
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = true
|
||||||
|
s.rawPkt = pkt
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
||||||
|
if info.hdrLen > tcpCoalesceHdrCap || info.hdrLen+info.payLen > tcpCoalesceBufSize {
|
||||||
|
// Pathological shape — can't fit our scratch, emit as-is.
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||||
|
s.hdrLen = info.hdrLen
|
||||||
|
s.ipHdrLen = info.ipHdrLen
|
||||||
|
s.isV6 = info.fk.isV6
|
||||||
|
s.fk = info.fk
|
||||||
|
s.gsoSize = info.payLen
|
||||||
|
s.numSeg = 1
|
||||||
|
s.totalPay = info.payLen
|
||||||
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
|
s.psh = info.flags&tcpFlagPsh != 0
|
||||||
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
if !s.psh {
|
||||||
|
c.openSlots[info.fk] = s
|
||||||
|
c.lastSlot = s
|
||||||
|
} else if last := c.lastSlot; last != nil && last.fk == info.fk {
|
||||||
|
// PSH-on-seed seals the slot immediately. Any prior cached open
|
||||||
|
// slot for this flow has just been sealed-and-replaced by this
|
||||||
|
// passthrough-shaped seed, so drop the cache too.
|
||||||
|
c.lastSlot = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// canAppend reports whether info's packet extends the slot's seed: same
|
||||||
|
// header shape and stable contents, adjacent seq, not oversized, chain not
|
||||||
|
// closed.
|
||||||
|
func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bool {
|
||||||
|
if s.psh {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.hdrLen != s.hdrLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.seq != s.nextSeq {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if s.numSeg >= tcpCoalesceMaxSegs {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.payLen > s.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// ECE state must be stable across a burst — receivers expect the
|
||||||
|
// flag set on every segment of a CE-echoing window or none.
|
||||||
|
seedFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||||
|
if (seedFlags^info.flags)&tcpFlagEce != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP) {
|
||||||
|
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
s.numSeg++
|
||||||
|
s.totalPay += info.payLen
|
||||||
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
|
if info.flags&tcpFlagPsh != 0 {
|
||||||
|
// Propagate PSH into the seed header so kernel TSO sets it on the
|
||||||
|
// last segment. Without this the sender's push signal is dropped.
|
||||||
|
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 {
|
||||||
|
s.psh = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) take() *coalesceSlot {
|
||||||
|
if n := len(c.pool); n > 0 {
|
||||||
|
s := c.pool[n-1]
|
||||||
|
c.pool[n-1] = nil
|
||||||
|
c.pool = c.pool[:n-1]
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return &coalesceSlot{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
clear(s.payIovs)
|
||||||
|
s.payIovs = s.payIovs[:0]
|
||||||
|
s.numSeg = 0
|
||||||
|
s.totalPay = 0
|
||||||
|
s.psh = false
|
||||||
|
c.pool = append(c.pool, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushSlot patches the header and calls WriteGSO. Does not remove the
|
||||||
|
// slot from c.slots.
|
||||||
|
func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
||||||
|
total := s.hdrLen + s.totalPay
|
||||||
|
l4Len := total - s.ipHdrLen
|
||||||
|
hdr := s.hdrBuf[:s.hdrLen]
|
||||||
|
|
||||||
|
if s.isV6 {
|
||||||
|
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||||
|
} else {
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||||
|
hdr[10] = 0
|
||||||
|
hdr[11] = 0
|
||||||
|
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||||
|
}
|
||||||
|
|
||||||
|
var psum uint32
|
||||||
|
if s.isV6 {
|
||||||
|
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoTCP, l4Len)
|
||||||
|
} else {
|
||||||
|
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoTCP, l4Len)
|
||||||
|
}
|
||||||
|
tcsum := s.ipHdrLen + 16
|
||||||
|
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||||
|
|
||||||
|
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
||||||
|
}
|
||||||
|
|
||||||
|
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||||
|
// equality on every field that must be identical across coalesced
|
||||||
|
// 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 {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !ipHeadersMatch(a, b, isV6) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
||||||
|
// [18:tcpHdrLen] options (incl. urgent).
|
||||||
|
tcp := ipHdrLen
|
||||||
|
if !bytes.Equal(a[tcp:tcp+4], b[tcp:tcp+4]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[tcp+8:tcp+13], b[tcp+8:tcp+13]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[tcp+14:tcp+16], b[tcp+14:tcp+16]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[tcp+18:], b[tcp+18:]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// reorderForFlush neutralizes wire-side reorder that the rxOrder buffer
|
||||||
|
// couldn't catch (anything crossing a recvmmsg batch boundary). Without
|
||||||
|
// this pass a small wire reorder — counter 250 arriving in batch K when
|
||||||
|
// 200..249 are coming in batch K+1 — would seed an out-of-seq slot first
|
||||||
|
// and emit it ahead of the lower-seq slot, manifesting at the inner TCP
|
||||||
|
// receiver as a much larger reorder than the wire actually had.
|
||||||
|
//
|
||||||
|
// Two phases:
|
||||||
|
// 1. Sort each passthrough-bounded segment of c.slots by (flow, seq).
|
||||||
|
// Cross-flow ordering inside a segment isn't preserved (it never was
|
||||||
|
// and doesn't matter for any single flow's TCP correctness).
|
||||||
|
// 2. Sweep once and merge adjacent same-flow slots whose ranges are now
|
||||||
|
// contiguous AND whose tail is gsoSize-aligned. The tail constraint
|
||||||
|
// matters because the kernel TSO splitter chops at gsoSize from the
|
||||||
|
// start of the merged payload — a short segment in the middle would
|
||||||
|
// desynchronize every later segment.
|
||||||
|
//
|
||||||
|
// Passthrough slots act as barriers: the merge check skips them on either
|
||||||
|
// side, so a SYN/FIN/RST/CWR is never reordered relative to its flow's
|
||||||
|
// data.
|
||||||
|
func (c *TCPCoalescer) reorderForFlush() {
|
||||||
|
if len(c.slots) <= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
runStart := 0
|
||||||
|
for i := 0; i <= len(c.slots); i++ {
|
||||||
|
if i < len(c.slots) && !c.slots[i].passthrough {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
c.sortRun(c.slots[runStart:i])
|
||||||
|
runStart = i + 1
|
||||||
|
}
|
||||||
|
out := c.slots[:0]
|
||||||
|
logged := false
|
||||||
|
for _, s := range c.slots {
|
||||||
|
if n := len(out); n > 0 {
|
||||||
|
prev := out[n-1]
|
||||||
|
if !prev.passthrough && !s.passthrough && prev.fk == s.fk {
|
||||||
|
// Same-flow neighbors after sort. If they aren't seq-
|
||||||
|
// contiguous it's a real gap — packets the wire reordered
|
||||||
|
// across batches, or actual loss before nebula. Log it so
|
||||||
|
// the operator can quantify how often it happens; the data
|
||||||
|
// itself still emits in seq order, kernel TCP handles the
|
||||||
|
// gap via its OOO queue.
|
||||||
|
if prev.nextSeq != slotSeedSeq(s) {
|
||||||
|
logged = true
|
||||||
|
gap := int64(slotSeedSeq(s)) - int64(prev.nextSeq)
|
||||||
|
slog.Default().Warn("tcp coalesce: cross-slot seq gap",
|
||||||
|
"src", flowKeyAddr(s.fk, false),
|
||||||
|
"dst", flowKeyAddr(s.fk, true),
|
||||||
|
"sport", s.fk.sport,
|
||||||
|
"dport", s.fk.dport,
|
||||||
|
"prev_seed_seq", slotSeedSeq(prev),
|
||||||
|
"prev_next_seq", prev.nextSeq,
|
||||||
|
"this_seed_seq", slotSeedSeq(s),
|
||||||
|
"gap_bytes", gap,
|
||||||
|
"prev_seg_count", prev.numSeg,
|
||||||
|
"prev_total_pay", prev.totalPay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if canMergeSlots(prev, s) {
|
||||||
|
mergeSlots(prev, s)
|
||||||
|
c.release(s)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out = append(out, s)
|
||||||
|
}
|
||||||
|
if logged {
|
||||||
|
slog.Default().Warn("==== end of batch ====")
|
||||||
|
}
|
||||||
|
c.slots = out
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyAddr returns the src or dst address from fk as a netip.Addr for
|
||||||
|
// logging. Only used on the cold gap-log path so the netip allocation
|
||||||
|
// doesn't matter.
|
||||||
|
func flowKeyAddr(fk flowKey, dst bool) netip.Addr {
|
||||||
|
src := fk.src
|
||||||
|
if dst {
|
||||||
|
src = fk.dst
|
||||||
|
}
|
||||||
|
if fk.isV6 {
|
||||||
|
return netip.AddrFrom16(src)
|
||||||
|
}
|
||||||
|
var v4 [4]byte
|
||||||
|
copy(v4[:], src[:4])
|
||||||
|
return netip.AddrFrom4(v4)
|
||||||
|
}
|
||||||
|
|
||||||
|
// sortRun stable-sorts run by (flowKey, seedSeq) so each flow's slots
|
||||||
|
// cluster together in seq order, ready for the merge sweep. Stable so
|
||||||
|
// equal-key slots keep their original relative position (defensive — a
|
||||||
|
// duplicate seedSeq would already mean something's wrong upstream).
|
||||||
|
func (c *TCPCoalescer) sortRun(run []*coalesceSlot) {
|
||||||
|
if len(run) <= 1 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// slices.SortStableFunc with a free, non-capturing comparator avoids the
|
||||||
|
// reflection + closure-escape allocations that sort.SliceStable forces.
|
||||||
|
slices.SortStableFunc(run, compareCoalesceSlots)
|
||||||
|
}
|
||||||
|
|
||||||
|
func compareCoalesceSlots(a, b *coalesceSlot) int {
|
||||||
|
if cmp := flowKeyCompare(a.fk, b.fk); cmp != 0 {
|
||||||
|
return cmp
|
||||||
|
}
|
||||||
|
aSeq, bSeq := slotSeedSeq(a), slotSeedSeq(b)
|
||||||
|
if aSeq == bSeq {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if tcpSeqLess(aSeq, bSeq) {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// slotSeedSeq returns the TCP seq of the slot's seed (first segment).
|
||||||
|
// nextSeq tracks the seq just past the last appended byte; subtracting
|
||||||
|
// totalPay walks back to the seed. uint32 wraparound is the right TCP
|
||||||
|
// arithmetic so no special-casing is needed.
|
||||||
|
func slotSeedSeq(s *coalesceSlot) uint32 {
|
||||||
|
return s.nextSeq - uint32(s.totalPay)
|
||||||
|
}
|
||||||
|
|
||||||
|
// tcpSeqLess reports whether a precedes b in TCP serial-number arithmetic
|
||||||
|
// (RFC 1323 §2.3). The signed int32 cast turns the modular subtraction
|
||||||
|
// into the right comparison even across the 2^32 wrap.
|
||||||
|
func tcpSeqLess(a, b uint32) bool {
|
||||||
|
return int32(a-b) < 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyCompare orders flowKeys deterministically. The exact ordering
|
||||||
|
// is irrelevant — only that same-flow slots cluster together so the
|
||||||
|
// post-sort sweep can merge contiguous pairs.
|
||||||
|
func flowKeyCompare(a, b flowKey) int {
|
||||||
|
// Cheap scalar fields first so most non-matching keys short-circuit
|
||||||
|
// without ever calling bytes.Compare. sport is the ephemeral port on
|
||||||
|
// egress flows and discriminates fastest. For matching keys (same
|
||||||
|
// flow), array equality on src/dst inlines to word-sized compares,
|
||||||
|
// so we only pay bytes.Compare when the arrays actually differ.
|
||||||
|
if a.sport != b.sport {
|
||||||
|
if a.sport < b.sport {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
if a.dport != b.dport {
|
||||||
|
if a.dport < b.dport {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
if a.dst != b.dst {
|
||||||
|
return bytes.Compare(a.dst[:], b.dst[:])
|
||||||
|
}
|
||||||
|
if a.src != b.src {
|
||||||
|
return bytes.Compare(a.src[:], b.src[:])
|
||||||
|
}
|
||||||
|
if a.isV6 != b.isV6 {
|
||||||
|
if !a.isV6 {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// canMergeSlots reports whether s can fold into prev as one merged TSO
|
||||||
|
// superpacket. Same flow, contiguous TCP byte range, equal gsoSize, and
|
||||||
|
// fits within the kernel TSO limits. The tail-of-prev check rejects any
|
||||||
|
// merge whose first slot ended on a sub-gsoSize segment — kernel TSO
|
||||||
|
// would split the merged skb at gsoSize boundaries from the start, so a
|
||||||
|
// short segment in the middle would corrupt every later segment. PSH and
|
||||||
|
// ECE state must agree across both slots: PSH is a semantic delimiter
|
||||||
|
// (preserving the sender's push boundary) and ECE state must be uniform
|
||||||
|
// across a window (the same rule canAppend enforces for in-flow appends).
|
||||||
|
//
|
||||||
|
// Note: a slot sealed by reorder (canAppend returned false on seq
|
||||||
|
// mismatch) keeps psh=false, so this restriction does not block the
|
||||||
|
// reorder-fix merge — only legitimate PSH-set seals.
|
||||||
|
func canMergeSlots(prev, s *coalesceSlot) bool {
|
||||||
|
if prev.psh {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.fk != s.fk {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.gsoSize != s.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.nextSeq != slotSeedSeq(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.numSeg+s.numSeg > tcpCoalesceMaxSegs {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if prev.hdrLen+prev.totalPay+s.totalPay > tcpCoalesceBufSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if len(prev.payIovs[len(prev.payIovs)-1]) != prev.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
prevFlags := prev.hdrBuf[prev.ipHdrLen+13]
|
||||||
|
sFlags := s.hdrBuf[s.ipHdrLen+13]
|
||||||
|
if (prevFlags^sFlags)&tcpFlagEce != 0 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !headersMatch(prev.hdrBuf[:prev.hdrLen], s.hdrBuf[:s.hdrLen], prev.isV6, prev.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// mergeSlots folds src into dst in place: payIovs concatenated, counters
|
||||||
|
// and totals updated, PSH and IP-level CE bits OR'd into the seed header
|
||||||
|
// so neither the push signal nor a CE mark is lost. The seed header's
|
||||||
|
// 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) {
|
||||||
|
dst.payIovs = append(dst.payIovs, src.payIovs...)
|
||||||
|
dst.numSeg += src.numSeg
|
||||||
|
dst.totalPay += src.totalPay
|
||||||
|
dst.nextSeq = src.nextSeq
|
||||||
|
if src.psh {
|
||||||
|
dst.psh = true
|
||||||
|
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
|
||||||
|
// already have its checksum field zeroed) and returns the folded/inverted
|
||||||
|
// 16-bit value to store.
|
||||||
|
func ipv4HdrChecksum(hdr []byte) uint16 {
|
||||||
|
var sum uint32
|
||||||
|
for i := 0; i+1 < len(hdr); i += 2 {
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(hdr[i : i+2]))
|
||||||
|
}
|
||||||
|
if len(hdr)%2 == 1 {
|
||||||
|
sum += uint32(hdr[len(hdr)-1]) << 8
|
||||||
|
}
|
||||||
|
for sum>>16 != 0 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
}
|
||||||
|
return ^uint16(sum)
|
||||||
|
}
|
||||||
|
|
||||||
|
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
|
||||||
|
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
||||||
|
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
|
||||||
|
// reuses these helpers.
|
||||||
|
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||||
|
var sum uint32
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(src[2:4]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(dst[0:2]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(dst[2:4]))
|
||||||
|
sum += uint32(proto)
|
||||||
|
sum += uint32(l4Len)
|
||||||
|
return sum
|
||||||
|
}
|
||||||
|
|
||||||
|
func pseudoSumIPv6(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||||
|
var sum uint32
|
||||||
|
for i := 0; i < 16; i += 2 {
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(src[i : i+2]))
|
||||||
|
sum += uint32(binary.BigEndian.Uint16(dst[i : i+2]))
|
||||||
|
}
|
||||||
|
sum += uint32(l4Len >> 16)
|
||||||
|
sum += uint32(l4Len & 0xffff)
|
||||||
|
sum += uint32(proto)
|
||||||
|
return sum
|
||||||
|
}
|
||||||
|
|
||||||
|
// foldOnceNoInvert folds the 32-bit accumulator to 16 bits and returns it
|
||||||
|
// unchanged (no one's complement). This is what virtio NEEDS_CSUM wants in
|
||||||
|
// the L4 checksum field — the kernel will add the payload sum and invert.
|
||||||
|
func foldOnceNoInvert(sum uint32) uint16 {
|
||||||
|
for sum>>16 != 0 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
}
|
||||||
|
return uint16(sum)
|
||||||
|
}
|
||||||
@@ -0,0 +1,239 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
|
)
|
||||||
|
|
||||||
|
// nopTunWriter is a zero-alloc tio.GSOWriter for benchmarks. Discards
|
||||||
|
// everything but satisfies the interface the coalescer detects.
|
||||||
|
type nopTunWriter struct{}
|
||||||
|
|
||||||
|
func (nopTunWriter) Write(p []byte) (int, error) { return len(p), nil }
|
||||||
|
func (nopTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (nopTunWriter) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{TSO: true, USO: true}
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTCPv4BulkFlow returns a slice of N adjacent ACK-only TCP segments
|
||||||
|
// on a single 5-tuple, each carrying payloadLen bytes. Seq numbers are
|
||||||
|
// contiguous so every packet is coalesceable onto the previous one.
|
||||||
|
func buildTCPv4BulkFlow(n, payloadLen int) [][]byte {
|
||||||
|
pkts := make([][]byte, n)
|
||||||
|
pay := make([]byte, payloadLen)
|
||||||
|
seq := uint32(1000)
|
||||||
|
for i := range n {
|
||||||
|
pkts[i] = buildTCPv4(seq, tcpAck, pay)
|
||||||
|
seq += uint32(payloadLen)
|
||||||
|
}
|
||||||
|
return pkts
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildTCPv4Interleaved returns nFlows * perFlow packets with per-flow
|
||||||
|
// seq continuity but round-robin across flows — worst case for any
|
||||||
|
// "last-slot" cache.
|
||||||
|
func buildTCPv4Interleaved(nFlows, perFlow, payloadLen int) [][]byte {
|
||||||
|
pay := make([]byte, payloadLen)
|
||||||
|
seqs := make([]uint32, nFlows)
|
||||||
|
for i := range seqs {
|
||||||
|
seqs[i] = uint32(1000 + i*1000000)
|
||||||
|
}
|
||||||
|
pkts := make([][]byte, 0, nFlows*perFlow)
|
||||||
|
for range perFlow {
|
||||||
|
for f := range nFlows {
|
||||||
|
sport := uint16(10000 + f)
|
||||||
|
pkts = append(pkts, buildTCPv4Ports(sport, 2000, seqs[f], tcpAck, pay))
|
||||||
|
seqs[f] += uint32(payloadLen)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return pkts
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildICMPv4 returns a minimal non-TCP packet that takes the passthrough
|
||||||
|
// branch in Commit.
|
||||||
|
func buildICMPv4() []byte {
|
||||||
|
pkt := make([]byte, 28)
|
||||||
|
pkt[0] = 0x45
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||||
|
pkt[9] = 1 // ICMP
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
// runCommitBench drives Commit over pkts batchSize at a time, flushing
|
||||||
|
// between batches, and reports per-packet cost.
|
||||||
|
func runCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
|
b.Helper()
|
||||||
|
c := NewTCPCoalescer(nopTunWriter{})
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
pkt := pkts[i%len(pkts)]
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
if (i+1)%batchSize == 0 {
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Drain any trailing partial batch so slot state doesn't leak across runs.
|
||||||
|
_ = c.Flush()
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitSingleFlow is the bulk-TCP steady state: one flow,
|
||||||
|
// contiguous seq, 1200-byte payloads. Every packet past the seed should
|
||||||
|
// append onto the open slot. This is the case we most care about.
|
||||||
|
func BenchmarkCommitSingleFlow(b *testing.B) {
|
||||||
|
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||||
|
runCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitInterleaved4 has 4 concurrent bulk flows round-robined.
|
||||||
|
// A single-entry fast-path cache will miss on every packet; an N-way
|
||||||
|
// cache or map lookup carries the weight.
|
||||||
|
func BenchmarkCommitInterleaved4(b *testing.B) {
|
||||||
|
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||||
|
runCommitBench(b, pkts, len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitInterleaved16 stresses the map at higher flow counts.
|
||||||
|
func BenchmarkCommitInterleaved16(b *testing.B) {
|
||||||
|
pkts := buildTCPv4Interleaved(16, tcpCoalesceMaxSegs, 1200)
|
||||||
|
runCommitBench(b, pkts, len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitPassthrough exercises the non-TCP branch: parseTCPBase
|
||||||
|
// bails early and addPassthrough is the only work.
|
||||||
|
func BenchmarkCommitPassthrough(b *testing.B) {
|
||||||
|
pkt := buildICMPv4()
|
||||||
|
pkts := make([][]byte, 64)
|
||||||
|
for i := range pkts {
|
||||||
|
pkts[i] = pkt
|
||||||
|
}
|
||||||
|
runCommitBench(b, pkts, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkCommitNonCoalesceableTCP sends SYN|ACK packets on one flow.
|
||||||
|
// Each packet takes the "TCP but not admissible" branch which does a
|
||||||
|
// map delete + passthrough. Measures the seal-without-slot cost.
|
||||||
|
func BenchmarkCommitNonCoalesceableTCP(b *testing.B) {
|
||||||
|
pay := make([]byte, 0)
|
||||||
|
pkts := make([][]byte, 64)
|
||||||
|
for i := range pkts {
|
||||||
|
pkts[i] = buildTCPv4(uint32(1000+i), tcpSyn|tcpAck, pay)
|
||||||
|
}
|
||||||
|
runCommitBench(b, pkts, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
// runMultiCommitBench drives MultiCoalescer.Commit. The dispatcher does
|
||||||
|
// the IP/L4 parse once and passes the parsed struct to the lane, so this
|
||||||
|
// is the bench that shows the savings of skipping the lane's re-parse.
|
||||||
|
func runMultiCommitBench(b *testing.B, pkts [][]byte, batchSize int) {
|
||||||
|
b.Helper()
|
||||||
|
m := NewMultiCoalescer(nopTunWriter{}, true, true)
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.SetBytes(int64(len(pkts[0])))
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
pkt := pkts[i%len(pkts)]
|
||||||
|
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()
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkMultiCommitSingleFlow is the multi-lane analogue of
|
||||||
|
// BenchmarkCommitSingleFlow — same workload but routed through the
|
||||||
|
// dispatcher. The delta vs the single-lane bench measures dispatcher
|
||||||
|
// overhead.
|
||||||
|
func BenchmarkMultiCommitSingleFlow(b *testing.B) {
|
||||||
|
pkts := buildTCPv4BulkFlow(tcpCoalesceMaxSegs, 1200)
|
||||||
|
runMultiCommitBench(b, pkts, tcpCoalesceMaxSegs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkMultiCommitInterleaved4 mirrors BenchmarkCommitInterleaved4
|
||||||
|
// through the dispatcher.
|
||||||
|
func BenchmarkMultiCommitInterleaved4(b *testing.B) {
|
||||||
|
pkts := buildTCPv4Interleaved(4, tcpCoalesceMaxSegs, 1200)
|
||||||
|
runMultiCommitBench(b, pkts, len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyPair is one comparison input for the flowKeyCompare bench.
|
||||||
|
type flowKeyPair struct{ a, b flowKey }
|
||||||
|
|
||||||
|
// makeFlowKey builds an IPv4 flowKey from compact inputs.
|
||||||
|
func makeFlowKey(srcLow, dstLow uint32, sport, dport uint16) flowKey {
|
||||||
|
var fk flowKey
|
||||||
|
binary.BigEndian.PutUint32(fk.src[12:16], srcLow)
|
||||||
|
binary.BigEndian.PutUint32(fk.dst[12:16], dstLow)
|
||||||
|
fk.sport = sport
|
||||||
|
fk.dport = dport
|
||||||
|
return fk
|
||||||
|
}
|
||||||
|
|
||||||
|
// flowKeyCases are the workload mixes flowKeyCompare sees in practice.
|
||||||
|
// - sameFlow: equal keys; tests the equal-path cost (sort runs hit this
|
||||||
|
// repeatedly when many segments share a flow).
|
||||||
|
// - sportDiffers: same src/dst/dport, different sport — the typical
|
||||||
|
// "sibling flows from one host to one server" pattern.
|
||||||
|
// - dstDiffers: same src/sport/dport, different dst — outbound to many
|
||||||
|
// servers from a fixed local port.
|
||||||
|
// - allDiffer: every field differs; worst case for short-circuiting.
|
||||||
|
func flowKeyCases() map[string][]flowKeyPair {
|
||||||
|
const n = 64
|
||||||
|
cases := map[string][]flowKeyPair{
|
||||||
|
"sameFlow": make([]flowKeyPair, n),
|
||||||
|
"sportDiffers": make([]flowKeyPair, n),
|
||||||
|
"dstDiffers": make([]flowKeyPair, n),
|
||||||
|
"allDiffer": make([]flowKeyPair, n),
|
||||||
|
}
|
||||||
|
for i := range n {
|
||||||
|
base := makeFlowKey(0x0a000001, 0x0a000002, 40000, 443)
|
||||||
|
cases["sameFlow"][i] = flowKeyPair{a: base, b: base}
|
||||||
|
cases["sportDiffers"][i] = flowKeyPair{
|
||||||
|
a: base,
|
||||||
|
b: makeFlowKey(0x0a000001, 0x0a000002, uint16(40001+i), 443),
|
||||||
|
}
|
||||||
|
cases["dstDiffers"][i] = flowKeyPair{
|
||||||
|
a: base,
|
||||||
|
b: makeFlowKey(0x0a000001, uint32(0x0a000002+i+1), 40000, 443),
|
||||||
|
}
|
||||||
|
cases["allDiffer"][i] = flowKeyPair{
|
||||||
|
a: makeFlowKey(uint32(0x0a000001+i), uint32(0x0a000002+i), uint16(40000+i), uint16(80+i)),
|
||||||
|
b: makeFlowKey(uint32(0x0b000001+i), uint32(0x0b000002+i), uint16(50000+i), uint16(443+i)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return cases
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkFlowKeyCompare measures flowKeyCompare across the workloads
|
||||||
|
// the sort step actually sees. Use this to compare reorderings.
|
||||||
|
func BenchmarkFlowKeyCompare(b *testing.B) {
|
||||||
|
for name, pairs := range flowKeyCases() {
|
||||||
|
b.Run(name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var sink int
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
p := pairs[i&(len(pairs)-1)]
|
||||||
|
sink += flowKeyCompare(p.a, p.b)
|
||||||
|
}
|
||||||
|
runtime.KeepAlive(sink)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,64 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import "net/netip"
|
||||||
|
|
||||||
|
const SendBatchCap = 128
|
||||||
|
|
||||||
|
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
||||||
|
type batchWriter interface {
|
||||||
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||||
|
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
||||||
|
// 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 {
|
||||||
|
out batchWriter
|
||||||
|
bufs [][]byte
|
||||||
|
dsts []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
backing []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewSendBatch(out batchWriter, batchCap, slotCap int) *SendBatch {
|
||||||
|
return &SendBatch{
|
||||||
|
out: out,
|
||||||
|
bufs: make([][]byte, 0, batchCap),
|
||||||
|
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||||
|
ecns: make([]byte, 0, batchCap),
|
||||||
|
backing: make([]byte, 0, batchCap*slotCap),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Reserve(sz int) []byte {
|
||||||
|
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]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
|
||||||
|
b.bufs = append(b.bufs, pkt)
|
||||||
|
b.dsts = append(b.dsts, dst)
|
||||||
|
b.ecns = append(b.ecns, outerECN)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Flush() error {
|
||||||
|
var err error
|
||||||
|
if len(b.bufs) > 0 {
|
||||||
|
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
||||||
|
}
|
||||||
|
clear(b.bufs)
|
||||||
|
b.bufs = b.bufs[:0]
|
||||||
|
b.dsts = b.dsts[:0]
|
||||||
|
b.ecns = b.ecns[:0]
|
||||||
|
b.backing = b.backing[:0]
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeBatchWriter struct {
|
||||||
|
bufs [][]byte
|
||||||
|
addrs []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||||
|
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||||
|
// returns, so tests must capture data before that happens.
|
||||||
|
w.bufs = make([][]byte, len(bufs))
|
||||||
|
for i, b := range bufs {
|
||||||
|
cp := make([]byte, len(b))
|
||||||
|
copy(cp, b)
|
||||||
|
w.bufs[i] = cp
|
||||||
|
}
|
||||||
|
w.addrs = append(w.addrs[:0], addrs...)
|
||||||
|
w.ecns = append(w.ecns[:0], ecns...)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
b := NewSendBatch(fw, 4, 32)
|
||||||
|
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
slot := b.Reserve(32)
|
||||||
|
if cap(slot) != 32 {
|
||||||
|
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||||
|
}
|
||||||
|
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||||
|
b.Commit(pkt, ap, 0)
|
||||||
|
}
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
if len(fw.bufs) != 4 {
|
||||||
|
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
||||||
|
}
|
||||||
|
for i, buf := range fw.bufs {
|
||||||
|
if len(buf) != 3 || buf[0] != byte(i) {
|
||||||
|
t.Errorf("buf %d: %x", i, buf)
|
||||||
|
}
|
||||||
|
if fw.addrs[i] != ap {
|
||||||
|
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush again with nothing committed — should be a no-op.
|
||||||
|
fw.bufs = nil
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("empty Flush: %v", err)
|
||||||
|
}
|
||||||
|
if fw.bufs != nil {
|
||||||
|
t.Fatalf("empty Flush triggered WriteBatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reuse after Flush.
|
||||||
|
slot := b.Reserve(32)
|
||||||
|
if cap(slot) != 32 {
|
||||||
|
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
b := NewSendBatch(fw, 3, 8)
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||||
|
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
s := b.Reserve(8)
|
||||||
|
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||||
|
b.Commit(pkt, ap, 0)
|
||||||
|
}
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, buf := range fw.bufs {
|
||||||
|
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
||||||
|
t.Errorf("slot %d corrupted: %x", i, buf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
// Tiny initial backing forces a grow on the second Reserve.
|
||||||
|
b := NewSendBatch(fw, 1, 4)
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||||
|
|
||||||
|
s1 := b.Reserve(4)
|
||||||
|
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
||||||
|
b.Commit(pkt1, ap, 0)
|
||||||
|
|
||||||
|
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
||||||
|
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
||||||
|
b.Commit(pkt2, ap, 0)
|
||||||
|
|
||||||
|
// pkt1 must still be intact even though backing reallocated.
|
||||||
|
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
||||||
|
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
if len(fw.bufs) != 2 {
|
||||||
|
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
||||||
|
}
|
||||||
|
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
||||||
|
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
||||||
|
}
|
||||||
|
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
||||||
|
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,336 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ipProtoUDP is the IANA protocol number for UDP.
|
||||||
|
const ipProtoUDP = 17
|
||||||
|
|
||||||
|
// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the
|
||||||
|
// kernel's gso_max_size; payloads beyond this are emitted as-is.
|
||||||
|
const udpCoalesceBufSize = 65535
|
||||||
|
|
||||||
|
// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO
|
||||||
|
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
||||||
|
const udpCoalesceMaxSegs = 64
|
||||||
|
|
||||||
|
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
|
||||||
|
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
|
||||||
|
const udpCoalesceHdrCap = 64
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
passthrough bool
|
||||||
|
rawPkt []byte // borrowed when passthrough
|
||||||
|
|
||||||
|
fk flowKey
|
||||||
|
hdrBuf [udpCoalesceHdrCap]byte
|
||||||
|
hdrLen int
|
||||||
|
ipHdrLen int
|
||||||
|
isV6 bool
|
||||||
|
gsoSize int // per-segment UDP payload length
|
||||||
|
numSeg int
|
||||||
|
totalPay int
|
||||||
|
// sealed closes the chain: set when a sub-gsoSize segment is appended
|
||||||
|
// (kernel UDP-GSO requires every segment but the last to be exactly
|
||||||
|
// gsoSize) or when limits are hit. No further appends after.
|
||||||
|
sealed bool
|
||||||
|
payIovs [][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
||||||
|
// concurrent flows and emits each flow's run as a single GSO_UDP_L4
|
||||||
|
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the
|
||||||
|
// underlying writer doesn't support USO.
|
||||||
|
//
|
||||||
|
// All output — coalesced or not — is deferred until Flush so per-flow
|
||||||
|
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
|
||||||
|
// across the TCP/UDP/passthrough split when this coalescer runs alongside
|
||||||
|
// others — see multi_coalesce.go. Per-flow order is preserved because a
|
||||||
|
// single 5-tuple only ever lands in one lane and each lane preserves its
|
||||||
|
// own slot order.
|
||||||
|
//
|
||||||
|
// Owns no locks; one coalescer per TUN write queue.
|
||||||
|
type UDPCoalescer struct {
|
||||||
|
plainW io.Writer
|
||||||
|
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
|
||||||
|
|
||||||
|
slots []*udpSlot
|
||||||
|
openSlots map[flowKey]*udpSlot
|
||||||
|
pool []*udpSlot
|
||||||
|
|
||||||
|
backing []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
|
||||||
|
// this when the underlying Queue's Capabilities advertise USO; otherwise
|
||||||
|
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
|
||||||
|
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
|
||||||
|
// plain Writes — same defensive shape as the TCP coalescer.
|
||||||
|
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
||||||
|
c := &UDPCoalescer{
|
||||||
|
plainW: w,
|
||||||
|
slots: make([]*udpSlot, 0, initialSlots),
|
||||||
|
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
||||||
|
pool: make([]*udpSlot, 0, initialSlots),
|
||||||
|
backing: make([]byte, 0, initialSlots*udpCoalesceBufSize),
|
||||||
|
}
|
||||||
|
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
|
||||||
|
c.gsoW = gw
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// parsedUDP holds the fields extracted from a single parse so later steps
|
||||||
|
// (admission, slot lookup, canAppend) don't re-walk the header.
|
||||||
|
type parsedUDP struct {
|
||||||
|
fk flowKey
|
||||||
|
ipHdrLen int
|
||||||
|
hdrLen int // ipHdrLen + 8
|
||||||
|
payLen int
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
|
||||||
|
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
|
||||||
|
// (IPv4 with options/fragmentation, IPv6 with extension headers).
|
||||||
|
func parseUDP(pkt []byte) (parsedUDP, bool) {
|
||||||
|
var p parsedUDP
|
||||||
|
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
|
||||||
|
if !ok {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
pkt = ip.pkt
|
||||||
|
p.fk = ip.fk
|
||||||
|
p.ipHdrLen = ip.ipHdrLen
|
||||||
|
|
||||||
|
if len(pkt) < p.ipHdrLen+8 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.hdrLen = p.ipHdrLen + 8
|
||||||
|
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
||||||
|
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
|
||||||
|
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
p.payLen = udpLen - 8
|
||||||
|
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
||||||
|
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
||||||
|
return p, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) Reserve(sz int) []byte {
|
||||||
|
return reserveFromBacking(&c.backing, sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
||||||
|
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
info, ok := parseUDP(pkt)
|
||||||
|
if !ok {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.commitParsed(pkt, info)
|
||||||
|
}
|
||||||
|
|
||||||
|
// commitParsed is the post-parse half of Commit. The caller must have
|
||||||
|
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
|
||||||
|
// avoid re-walking the IP/UDP header.
|
||||||
|
func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
|
||||||
|
if c.gsoW == nil {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if open := c.openSlots[info.fk]; open != nil {
|
||||||
|
if c.canAppend(open, pkt, info) {
|
||||||
|
c.appendPayload(open, pkt, info)
|
||||||
|
if open.sealed {
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||||
|
delete(c.openSlots, info.fk)
|
||||||
|
}
|
||||||
|
c.seed(pkt, info)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) Flush() error {
|
||||||
|
var first error
|
||||||
|
for _, s := range c.slots {
|
||||||
|
var err error
|
||||||
|
if s.passthrough {
|
||||||
|
_, err = c.plainW.Write(s.rawPkt)
|
||||||
|
} else {
|
||||||
|
err = c.flushSlot(s)
|
||||||
|
}
|
||||||
|
if err != nil && first == nil {
|
||||||
|
first = err
|
||||||
|
}
|
||||||
|
c.release(s)
|
||||||
|
}
|
||||||
|
clear(c.slots)
|
||||||
|
c.slots = c.slots[:0]
|
||||||
|
clear(c.openSlots)
|
||||||
|
c.backing = c.backing[:0]
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = true
|
||||||
|
s.rawPkt = pkt
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
||||||
|
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
|
||||||
|
c.addPassthrough(pkt)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s := c.take()
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
||||||
|
s.hdrLen = info.hdrLen
|
||||||
|
s.ipHdrLen = info.ipHdrLen
|
||||||
|
s.isV6 = info.fk.isV6
|
||||||
|
s.fk = info.fk
|
||||||
|
s.gsoSize = info.payLen
|
||||||
|
s.numSeg = 1
|
||||||
|
s.totalPay = info.payLen
|
||||||
|
s.sealed = false
|
||||||
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
c.slots = append(c.slots, s)
|
||||||
|
c.openSlots[info.fk] = s
|
||||||
|
}
|
||||||
|
|
||||||
|
// canAppend reports whether info's packet extends the slot's seed.
|
||||||
|
// Kernel UDP-GSO requires every segment except possibly the last to be
|
||||||
|
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
||||||
|
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
|
||||||
|
if s.sealed {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.hdrLen != s.hdrLen {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if s.numSeg >= udpCoalesceMaxSegs {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if info.payLen > s.gsoSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
|
||||||
|
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
|
s.numSeg++
|
||||||
|
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 {
|
||||||
|
// Last-segment-can-be-shorter: this seals the chain.
|
||||||
|
s.sealed = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) take() *udpSlot {
|
||||||
|
if n := len(c.pool); n > 0 {
|
||||||
|
s := c.pool[n-1]
|
||||||
|
c.pool[n-1] = nil
|
||||||
|
c.pool = c.pool[:n-1]
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return &udpSlot{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCoalescer) release(s *udpSlot) {
|
||||||
|
s.passthrough = false
|
||||||
|
s.rawPkt = nil
|
||||||
|
clear(s.payIovs)
|
||||||
|
s.payIovs = s.payIovs[:0]
|
||||||
|
s.numSeg = 0
|
||||||
|
s.totalPay = 0
|
||||||
|
s.sealed = false
|
||||||
|
c.pool = append(c.pool, s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushSlot patches the IP header total length / IPv6 payload length and
|
||||||
|
// the UDP length to the *total* across all coalesced segments, then seeds
|
||||||
|
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
||||||
|
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and
|
||||||
|
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
|
||||||
|
// values would silently drop everything but the first segment. The kernel
|
||||||
|
// then walks each segment in __udp_gso_segment, recomputing per-segment
|
||||||
|
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
|
||||||
|
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
|
||||||
|
// our seed's uh->check must be consistent with the seed's uh->len, which
|
||||||
|
// is what passing the total to both pseudoSum and the UDP length field
|
||||||
|
// guarantees.
|
||||||
|
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
||||||
|
hdr := s.hdrBuf[:s.hdrLen]
|
||||||
|
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
||||||
|
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
||||||
|
|
||||||
|
if s.isV6 {
|
||||||
|
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
||||||
|
} else {
|
||||||
|
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
||||||
|
hdr[10] = 0
|
||||||
|
hdr[11] = 0
|
||||||
|
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
||||||
|
}
|
||||||
|
|
||||||
|
// UDP length field (offset 4 inside the UDP header) = total UDP size.
|
||||||
|
binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len))
|
||||||
|
|
||||||
|
var psum uint32
|
||||||
|
if s.isV6 {
|
||||||
|
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len)
|
||||||
|
} else {
|
||||||
|
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len)
|
||||||
|
}
|
||||||
|
udpCsumOff := s.ipHdrLen + 6
|
||||||
|
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
||||||
|
|
||||||
|
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
||||||
|
}
|
||||||
|
|
||||||
|
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
||||||
|
// every field that must be identical across coalesced segments. Length
|
||||||
|
// fields and the ECN bits in IP TOS/TC are masked out — appendPayload
|
||||||
|
// merges CE into the seed; flushSlot rewrites lengths.
|
||||||
|
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||||
|
if len(a) != len(b) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !ipHeadersMatch(a, b, isV6) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8] —
|
||||||
|
// length varies (we rewrite at flush) and the checksum will be redone.
|
||||||
|
udp := ipHdrLen
|
||||||
|
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -0,0 +1,383 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports.
|
||||||
|
func buildUDPv4(sport, dport uint16, payload []byte) []byte {
|
||||||
|
const ipHdrLen = 20
|
||||||
|
const udpHdrLen = 8
|
||||||
|
total := ipHdrLen + udpHdrLen + len(payload)
|
||||||
|
pkt := make([]byte, total)
|
||||||
|
|
||||||
|
pkt[0] = 0x45
|
||||||
|
pkt[1] = 0x00
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], 0)
|
||||||
|
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
|
||||||
|
pkt[8] = 64
|
||||||
|
pkt[9] = ipProtoUDP
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
|
||||||
|
binary.BigEndian.PutUint16(pkt[20:22], sport)
|
||||||
|
binary.BigEndian.PutUint16(pkt[22:24], dport)
|
||||||
|
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload)))
|
||||||
|
binary.BigEndian.PutUint16(pkt[26:28], 0)
|
||||||
|
|
||||||
|
copy(pkt[28:], payload)
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildUDPv6 builds a minimal IPv6+UDP packet.
|
||||||
|
func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
||||||
|
const ipHdrLen = 40
|
||||||
|
const udpHdrLen = 8
|
||||||
|
total := ipHdrLen + udpHdrLen + len(payload)
|
||||||
|
pkt := make([]byte, total)
|
||||||
|
|
||||||
|
pkt[0] = 0x60
|
||||||
|
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload)))
|
||||||
|
pkt[6] = ipProtoUDP
|
||||||
|
pkt[7] = 64
|
||||||
|
pkt[8] = 0xfe
|
||||||
|
pkt[9] = 0x80
|
||||||
|
pkt[23] = 1
|
||||||
|
pkt[24] = 0xfe
|
||||||
|
pkt[25] = 0x80
|
||||||
|
pkt[39] = 2
|
||||||
|
|
||||||
|
binary.BigEndian.PutUint16(pkt[40:42], sport)
|
||||||
|
binary.BigEndian.PutUint16(pkt[42:44], dport)
|
||||||
|
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload)))
|
||||||
|
binary.BigEndian.PutUint16(pkt[46:48], 0)
|
||||||
|
|
||||||
|
copy(pkt[48:], payload)
|
||||||
|
return pkt
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: false}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 100))
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
||||||
|
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
||||||
|
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
// ICMP packet
|
||||||
|
pkt := make([]byte, 28)
|
||||||
|
pkt[0] = 0x45
|
||||||
|
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
||||||
|
pkt[9] = 1
|
||||||
|
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
||||||
|
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
||||||
|
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("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
||||||
|
if err := c.Commit(pkt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
|
||||||
|
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
|
||||||
|
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||||
|
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 1 {
|
||||||
|
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||||
|
}
|
||||||
|
g := w.gsoWrites[0]
|
||||||
|
if g.gsoSize != 1200 {
|
||||||
|
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
|
||||||
|
}
|
||||||
|
if len(g.pays) != 3 {
|
||||||
|
t.Errorf("pay count=%d want 3", len(g.pays))
|
||||||
|
}
|
||||||
|
if g.csumStart != 20 {
|
||||||
|
t.Errorf("csumStart=%d want 20", g.csumStart)
|
||||||
|
}
|
||||||
|
// IP totalLen and UDP length must be the TOTAL across all segments —
|
||||||
|
// the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment
|
||||||
|
// value would silently drop everything but the first segment. Total =
|
||||||
|
// IP(20) + UDP(8) + 3*1200 = 3628.
|
||||||
|
gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4])
|
||||||
|
if gotTotalLen != 3628 {
|
||||||
|
t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen)
|
||||||
|
}
|
||||||
|
gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6])
|
||||||
|
if gotUDPLen != 8+3*1200 {
|
||||||
|
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Last segment may be shorter, sealing the chain.
|
||||||
|
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
full := make([]byte, 1200)
|
||||||
|
tail := make([]byte, 600)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// A 4th packet, even same-sized, must NOT join — chain is sealed.
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites[0].pays) != 3 {
|
||||||
|
t.Errorf("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites[1].pays) != 1 {
|
||||||
|
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
||||||
|
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Different 5-tuples must not coalesce.
|
||||||
|
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pay := make([]byte, 800)
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// Two flows × 2 datagrams each = 2 superpackets of 2 segments.
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
for i, g := range w.gsoWrites {
|
||||||
|
if len(g.pays) != 2 {
|
||||||
|
t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Caps at udpCoalesceMaxSegs.
|
||||||
|
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pay := make([]byte, 100)
|
||||||
|
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
||||||
|
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
// First superpacket holds udpCoalesceMaxSegs segments; the spillover
|
||||||
|
// reseeds a new one.
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs {
|
||||||
|
t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites[1].pays) != 5 {
|
||||||
|
t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CE marks on appended segments must be merged into the seed's IP TOS.
|
||||||
|
func TestUDPCoalescerMergesCEMark(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pay := make([]byte, 800)
|
||||||
|
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00
|
||||||
|
pkt1 := buildUDPv4(1000, 53, pay)
|
||||||
|
pkt1[1] = 0x03 // CE
|
||||||
|
pkt2 := buildUDPv4(1000, 53, pay)
|
||||||
|
if err := c.Commit(pkt0); err != nil {
|
||||||
|
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 {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 1 {
|
||||||
|
t.Fatalf("want 1 merged gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
||||||
|
}
|
||||||
|
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.
|
||||||
|
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pay := make([]byte, 1200)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 1 {
|
||||||
|
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
g := w.gsoWrites[0]
|
||||||
|
if !g.isV6 {
|
||||||
|
t.Errorf("expected v6 write")
|
||||||
|
}
|
||||||
|
if g.csumStart != 40 {
|
||||||
|
t.Errorf("csumStart=%d want 40", g.csumStart)
|
||||||
|
}
|
||||||
|
// IPv6 payload_len and UDP length must be TOTAL — kernel's
|
||||||
|
// ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 +
|
||||||
|
// 3*1200 = 3608.
|
||||||
|
gotPlen := binary.BigEndian.Uint16(g.hdr[4:6])
|
||||||
|
if gotPlen != 8+3*1200 {
|
||||||
|
t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200)
|
||||||
|
}
|
||||||
|
gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6])
|
||||||
|
if gotUDPLen != 8+3*1200 {
|
||||||
|
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// DSCP differences must reseed (headers don't match outside ECN).
|
||||||
|
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pay := make([]byte, 800)
|
||||||
|
pkt0 := buildUDPv4(1000, 53, pay)
|
||||||
|
pkt1 := buildUDPv4(1000, 53, pay)
|
||||||
|
pkt1[1] = 0xb8 // EF DSCP, ECN=0
|
||||||
|
if err := c.Commit(pkt0); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Commit(pkt1); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := c.Flush(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(w.gsoWrites) != 2 {
|
||||||
|
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fragmented IPv4 must not be coalesced.
|
||||||
|
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||||
|
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
||||||
|
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("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IPv4 with options is not admissible (we require IHL=5).
|
||||||
|
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
||||||
|
w := &fakeTunWriter{gsoEnabled: true}
|
||||||
|
c := NewUDPCoalescer(w)
|
||||||
|
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
||||||
|
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
||||||
|
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("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,23 @@
|
|||||||
|
package checksum
|
||||||
|
|
||||||
|
import (
|
||||||
|
"golang.org/x/sys/cpu"
|
||||||
|
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
)
|
||||||
|
|
||||||
|
//go:noescape
|
||||||
|
func checksumAVX2(buf []byte, initial uint16) uint16
|
||||||
|
|
||||||
|
var hasAVX2 = cpu.X86.HasAVX2
|
||||||
|
|
||||||
|
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||||
|
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
|
||||||
|
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
|
||||||
|
// falling back to gvisor's pure-Go implementation otherwise. The result
|
||||||
|
// matches gvisor's bit-for-bit for any buffer length and initial seed.
|
||||||
|
func Checksum(buf []byte, initial uint16) uint16 {
|
||||||
|
if hasAVX2 {
|
||||||
|
return checksumAVX2(buf, initial)
|
||||||
|
}
|
||||||
|
return gvisorchecksum.Checksum(buf, initial)
|
||||||
|
}
|
||||||
@@ -0,0 +1,157 @@
|
|||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
// func checksumAVX2(buf []byte, initial uint16) uint16
|
||||||
|
//
|
||||||
|
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
|
||||||
|
//
|
||||||
|
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
|
||||||
|
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
|
||||||
|
// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of
|
||||||
|
// headroom which is far more than the 16 KB/64 KB max practical inputs).
|
||||||
|
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
|
||||||
|
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
|
||||||
|
// equivalent to summing as 16-bit big-endian words.
|
||||||
|
//
|
||||||
|
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
|
||||||
|
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
|
||||||
|
// chunks, zero-extending each chunk's four uint32s into a ymm via
|
||||||
|
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
|
||||||
|
// chunk to break the dep chain. After the vector loop the lane sums are
|
||||||
|
// horizontally reduced and merged with a scalar accumulator that handles
|
||||||
|
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
|
||||||
|
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
|
||||||
|
MOVQ buf_base+0(FP), SI
|
||||||
|
MOVQ buf_len+8(FP), CX
|
||||||
|
MOVWQZX initial+24(FP), AX
|
||||||
|
|
||||||
|
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||||
|
// with the rest of the accumulator. The final fold's bswap16 will undo
|
||||||
|
// this and convert the whole result back to BE.
|
||||||
|
XCHGB AH, AL
|
||||||
|
|
||||||
|
CMPQ CX, $32
|
||||||
|
JLT scalar_tail
|
||||||
|
|
||||||
|
VPXOR Y4, Y4, Y4
|
||||||
|
VPXOR Y5, Y5, Y5
|
||||||
|
VPXOR Y6, Y6, Y6
|
||||||
|
VPXOR Y7, Y7, Y7
|
||||||
|
|
||||||
|
CMPQ CX, $64
|
||||||
|
JLT loop32
|
||||||
|
|
||||||
|
loop64:
|
||||||
|
VPMOVZXDQ (SI), Y0
|
||||||
|
VPMOVZXDQ 16(SI), Y1
|
||||||
|
VPMOVZXDQ 32(SI), Y2
|
||||||
|
VPMOVZXDQ 48(SI), Y3
|
||||||
|
VPADDQ Y0, Y4, Y4
|
||||||
|
VPADDQ Y1, Y5, Y5
|
||||||
|
VPADDQ Y2, Y6, Y6
|
||||||
|
VPADDQ Y3, Y7, Y7
|
||||||
|
ADDQ $64, SI
|
||||||
|
SUBQ $64, CX
|
||||||
|
CMPQ CX, $64
|
||||||
|
JGE loop64
|
||||||
|
|
||||||
|
loop32:
|
||||||
|
CMPQ CX, $32
|
||||||
|
JLT reduce_vec
|
||||||
|
VPMOVZXDQ (SI), Y0
|
||||||
|
VPMOVZXDQ 16(SI), Y1
|
||||||
|
VPADDQ Y0, Y4, Y4
|
||||||
|
VPADDQ Y1, Y5, Y5
|
||||||
|
ADDQ $32, SI
|
||||||
|
SUBQ $32, CX
|
||||||
|
JMP loop32
|
||||||
|
|
||||||
|
reduce_vec:
|
||||||
|
// Combine the four ymm accumulators into Y4.
|
||||||
|
VPADDQ Y5, Y4, Y4
|
||||||
|
VPADDQ Y7, Y6, Y6
|
||||||
|
VPADDQ Y6, Y4, Y4
|
||||||
|
|
||||||
|
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
|
||||||
|
VEXTRACTI128 $1, Y4, X5
|
||||||
|
VPADDQ X5, X4, X4
|
||||||
|
VPSHUFD $0x4e, X4, X5
|
||||||
|
VPADDQ X5, X4, X4
|
||||||
|
VMOVQ X4, R8
|
||||||
|
VZEROUPPER
|
||||||
|
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
|
||||||
|
scalar_tail:
|
||||||
|
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
|
||||||
|
CMPQ CX, $8
|
||||||
|
JLT tail4
|
||||||
|
|
||||||
|
loop8:
|
||||||
|
ADDQ (SI), AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
ADDQ $8, SI
|
||||||
|
SUBQ $8, CX
|
||||||
|
CMPQ CX, $8
|
||||||
|
JGE loop8
|
||||||
|
|
||||||
|
tail4:
|
||||||
|
CMPQ CX, $4
|
||||||
|
JLT tail2
|
||||||
|
MOVL (SI), R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
ADDQ $4, SI
|
||||||
|
SUBQ $4, CX
|
||||||
|
|
||||||
|
tail2:
|
||||||
|
CMPQ CX, $2
|
||||||
|
JLT tail1
|
||||||
|
MOVWQZX (SI), R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
ADDQ $2, SI
|
||||||
|
SUBQ $2, CX
|
||||||
|
|
||||||
|
tail1:
|
||||||
|
TESTQ CX, CX
|
||||||
|
JZ fold
|
||||||
|
MOVBQZX (SI), R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
ADCQ $0, AX
|
||||||
|
|
||||||
|
fold:
|
||||||
|
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
|
||||||
|
// gvisor's reduce(). Each pair (split, add) halves the live width;
|
||||||
|
// the truncation steps absorb the single bit that may be left over
|
||||||
|
// after each add so the next round's bound holds.
|
||||||
|
|
||||||
|
// 64 → 33 bits.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $32, R8
|
||||||
|
MOVL AX, AX
|
||||||
|
ADDQ R8, AX
|
||||||
|
|
||||||
|
// 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $32, R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
MOVL AX, AX
|
||||||
|
|
||||||
|
// 32 → 17 bits.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $16, R8
|
||||||
|
MOVWQZX AX, AX
|
||||||
|
ADDQ R8, AX
|
||||||
|
|
||||||
|
// 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
|
||||||
|
MOVQ AX, R8
|
||||||
|
SHRQ $16, R8
|
||||||
|
ADDQ R8, AX
|
||||||
|
|
||||||
|
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||||
|
// to big-endian to match the gvisor API contract.
|
||||||
|
XCHGB AH, AL
|
||||||
|
|
||||||
|
MOVW AX, ret+32(FP)
|
||||||
|
RET
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
package checksum
|
||||||
|
|
||||||
|
//go:noescape
|
||||||
|
func checksumNEON(buf []byte, initial uint16) uint16
|
||||||
|
|
||||||
|
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
||||||
|
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
|
||||||
|
// that dispatches to a hand-written NEON routine. NEON is mandatory in
|
||||||
|
// armv8 so no feature check is needed.
|
||||||
|
func Checksum(buf []byte, initial uint16) uint16 {
|
||||||
|
return checksumNEON(buf, initial)
|
||||||
|
}
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
#include "textflag.h"
|
||||||
|
|
||||||
|
// func checksumNEON(buf []byte, initial uint16) uint16
|
||||||
|
//
|
||||||
|
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
|
||||||
|
// a stream of uint32s in machine (little-endian) byte order, accumulating
|
||||||
|
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
|
||||||
|
// at the very end to recover the on-wire (big-endian) result.
|
||||||
|
//
|
||||||
|
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
|
||||||
|
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
|
||||||
|
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
|
||||||
|
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
|
||||||
|
// updated twice per iter to break the dep chain. Tail bytes go through a
|
||||||
|
// scalar ADCS chain seeded with the byte-swapped initial.
|
||||||
|
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
|
||||||
|
MOVD buf_base+0(FP), R0
|
||||||
|
MOVD buf_len+8(FP), R1
|
||||||
|
MOVHU initial+24(FP), R2
|
||||||
|
|
||||||
|
// Pre-byteswap initial into the LE-summing space so it merges directly
|
||||||
|
// with the rest of the accumulator.
|
||||||
|
REV16W R2, R2
|
||||||
|
|
||||||
|
MOVD ZR, R3 // scalar accumulator
|
||||||
|
|
||||||
|
CMP $32, R1
|
||||||
|
BLT scalar_tail
|
||||||
|
|
||||||
|
VEOR V8.B16, V8.B16, V8.B16
|
||||||
|
VEOR V9.B16, V9.B16, V9.B16
|
||||||
|
VEOR V10.B16, V10.B16, V10.B16
|
||||||
|
VEOR V11.B16, V11.B16, V11.B16
|
||||||
|
|
||||||
|
CMP $64, R1
|
||||||
|
BLT loop16_init
|
||||||
|
|
||||||
|
loop64:
|
||||||
|
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
|
||||||
|
VUADDW V0.S2, V8.D2, V8.D2
|
||||||
|
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||||
|
VUADDW V1.S2, V10.D2, V10.D2
|
||||||
|
VUADDW2 V1.S4, V11.D2, V11.D2
|
||||||
|
VUADDW V2.S2, V8.D2, V8.D2
|
||||||
|
VUADDW2 V2.S4, V9.D2, V9.D2
|
||||||
|
VUADDW V3.S2, V10.D2, V10.D2
|
||||||
|
VUADDW2 V3.S4, V11.D2, V11.D2
|
||||||
|
SUB $64, R1, R1
|
||||||
|
CMP $64, R1
|
||||||
|
BGE loop64
|
||||||
|
|
||||||
|
loop16_init:
|
||||||
|
CMP $16, R1
|
||||||
|
BLT reduce_vec
|
||||||
|
|
||||||
|
loop16:
|
||||||
|
VLD1.P 16(R0), [V0.B16]
|
||||||
|
VUADDW V0.S2, V8.D2, V8.D2
|
||||||
|
VUADDW2 V0.S4, V9.D2, V9.D2
|
||||||
|
SUB $16, R1, R1
|
||||||
|
CMP $16, R1
|
||||||
|
BGE loop16
|
||||||
|
|
||||||
|
reduce_vec:
|
||||||
|
// Combine the four accumulators into V8.
|
||||||
|
VADD V9.D2, V8.D2, V8.D2
|
||||||
|
VADD V11.D2, V10.D2, V10.D2
|
||||||
|
VADD V10.D2, V8.D2, V8.D2
|
||||||
|
|
||||||
|
// Horizontal-add the two lanes of V8.D2 into a single uint64.
|
||||||
|
VADDP V8.D2, V8.D2, V8.D2
|
||||||
|
VMOV V8.D[0], R8
|
||||||
|
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
|
||||||
|
scalar_tail:
|
||||||
|
CMP $8, R1
|
||||||
|
BLT tail4
|
||||||
|
|
||||||
|
loop8:
|
||||||
|
MOVD.P 8(R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
SUB $8, R1, R1
|
||||||
|
CMP $8, R1
|
||||||
|
BGE loop8
|
||||||
|
|
||||||
|
tail4:
|
||||||
|
CMP $4, R1
|
||||||
|
BLT tail2
|
||||||
|
MOVWU.P 4(R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
SUB $4, R1, R1
|
||||||
|
|
||||||
|
tail2:
|
||||||
|
CMP $2, R1
|
||||||
|
BLT tail1
|
||||||
|
MOVHU.P 2(R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
SUB $2, R1, R1
|
||||||
|
|
||||||
|
tail1:
|
||||||
|
CBZ R1, fold
|
||||||
|
MOVBU (R0), R8
|
||||||
|
ADDS R8, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
|
||||||
|
fold:
|
||||||
|
// Merge the byte-swapped initial into our LE-form accumulator.
|
||||||
|
ADDS R2, R3, R3
|
||||||
|
ADC ZR, R3, R3
|
||||||
|
|
||||||
|
// 64 → 33 bits.
|
||||||
|
LSR $32, R3, R8
|
||||||
|
AND $0xffffffff, R3, R3
|
||||||
|
ADD R8, R3, R3
|
||||||
|
|
||||||
|
// 33 → 32 (truncate after adding bit 32 back).
|
||||||
|
LSR $32, R3, R8
|
||||||
|
ADD R8, R3, R3
|
||||||
|
AND $0xffffffff, R3, R3
|
||||||
|
|
||||||
|
// 32 → 17.
|
||||||
|
LSR $16, R3, R8
|
||||||
|
AND $0xffff, R3, R3
|
||||||
|
ADD R8, R3, R3
|
||||||
|
|
||||||
|
// 17 → 16 (truncation absorbs bit 16 below).
|
||||||
|
LSR $16, R3, R8
|
||||||
|
ADD R8, R3, R3
|
||||||
|
|
||||||
|
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
||||||
|
// to big-endian to match the gvisor API contract. REV16W swaps bytes
|
||||||
|
// within each 16-bit halfword of the low 32 bits, so it acts as a
|
||||||
|
// 16-bit byte-swap on the live low 16.
|
||||||
|
REV16W R3, R3
|
||||||
|
AND $0xffff, R3, R3
|
||||||
|
|
||||||
|
MOVH R3, ret+32(FP)
|
||||||
|
RET
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
//go:build !amd64 && !arm64
|
||||||
|
|
||||||
|
package checksum
|
||||||
|
|
||||||
|
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
|
||||||
|
// Checksum delegates to gvisor on architectures without a hand-written body.
|
||||||
|
func Checksum(buf []byte, initial uint16) uint16 {
|
||||||
|
return gvisorchecksum.Checksum(buf, initial)
|
||||||
|
}
|
||||||
@@ -0,0 +1,190 @@
|
|||||||
|
package checksum
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"math/rand/v2"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
)
|
||||||
|
|
||||||
|
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
||||||
|
// seeds and a handful of starting alignments, asserting that our local
|
||||||
|
// Checksum matches gvisor's reference bit-for-bit.
|
||||||
|
func TestChecksumMatchesGvisor(t *testing.T) {
|
||||||
|
rng := rand.New(rand.NewPCG(1, 2))
|
||||||
|
const padFront = 16
|
||||||
|
|
||||||
|
// Random pool large enough for the longest case + alignment slop.
|
||||||
|
pool := make([]byte, 4096+padFront)
|
||||||
|
for i := range pool {
|
||||||
|
pool[i] = byte(rng.Uint32())
|
||||||
|
}
|
||||||
|
|
||||||
|
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
||||||
|
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
||||||
|
|
||||||
|
for length := 0; length <= 4096; length++ {
|
||||||
|
for _, seed := range seeds {
|
||||||
|
for _, off := range offsets {
|
||||||
|
if off+length > len(pool) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
buf := pool[off : off+length]
|
||||||
|
want := gvisorchecksum.Checksum(buf, seed)
|
||||||
|
got := Checksum(buf, seed)
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
||||||
|
length, off, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestChecksumPatternedBuffers exercises specific byte patterns that have
|
||||||
|
// historically tripped up checksum implementations: all-zero, all-0xff,
|
||||||
|
// alternating, and ascending sequences.
|
||||||
|
func TestChecksumPatternedBuffers(t *testing.T) {
|
||||||
|
for length := 0; length <= 256; length++ {
|
||||||
|
patterns := map[string][]byte{
|
||||||
|
"zeros": make([]byte, length),
|
||||||
|
"ones": bytes(length, 0xff),
|
||||||
|
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
||||||
|
"ascending": ascending(length),
|
||||||
|
}
|
||||||
|
for name, buf := range patterns {
|
||||||
|
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
||||||
|
want := gvisorchecksum.Checksum(buf, seed)
|
||||||
|
got := Checksum(buf, seed)
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
||||||
|
name, length, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func bytes(n int, v byte) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = v
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func pattern(n int, p []byte) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = p[i%len(p)]
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func ascending(n int) []byte {
|
||||||
|
b := make([]byte, n)
|
||||||
|
for i := range b {
|
||||||
|
b[i] = byte(i)
|
||||||
|
}
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestChecksumTailPaths targets every combination of (SIMD body iterations,
|
||||||
|
// trailing tail bytes) the asm handlers walk through. The tail handlers
|
||||||
|
// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by
|
||||||
|
// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a
|
||||||
|
// representative spread of k values, including k=0 (no main loop, all tail)
|
||||||
|
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
||||||
|
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
||||||
|
func TestChecksumTailPaths(t *testing.T) {
|
||||||
|
rng := rand.New(rand.NewPCG(42, 17))
|
||||||
|
const padFront = 16
|
||||||
|
const maxK = 8
|
||||||
|
|
||||||
|
pool := make([]byte, 64*maxK+padFront+64)
|
||||||
|
for i := range pool {
|
||||||
|
pool[i] = byte(rng.Uint32())
|
||||||
|
}
|
||||||
|
|
||||||
|
seeds := []uint16{0, 0xffff, 0xabcd}
|
||||||
|
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
||||||
|
|
||||||
|
for k := 0; k <= maxK; k++ {
|
||||||
|
for tail := 0; tail < 64; tail++ {
|
||||||
|
length := 64*k + tail
|
||||||
|
for _, seed := range seeds {
|
||||||
|
for _, off := range offsets {
|
||||||
|
if off+length > len(pool) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
buf := pool[off : off+length]
|
||||||
|
want := gvisorchecksum.Checksum(buf, seed)
|
||||||
|
got := Checksum(buf, seed)
|
||||||
|
if got != want {
|
||||||
|
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
||||||
|
k, tail, length, off, seed, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
|
||||||
|
// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail
|
||||||
|
// handler is meaningfully on the hot path. Sizes are picked to either exercise
|
||||||
|
// every tail branch (tiny lengths) or sit slightly off realistic packet
|
||||||
|
// boundaries (e.g. 1499 = MTU − 1).
|
||||||
|
func BenchmarkChecksumTailSizes(b *testing.B) {
|
||||||
|
sizes := []int{
|
||||||
|
1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail
|
||||||
|
33, 35, 47, 63, // one loop32 + assorted tails
|
||||||
|
65, 95, 127, // one loop64 + assorted tails
|
||||||
|
1447, 1471, 1499, 1501, // around MTU
|
||||||
|
8191, 8193, // around USO
|
||||||
|
65531, 65533, // near the kernel max
|
||||||
|
}
|
||||||
|
for _, size := range sizes {
|
||||||
|
buf := make([]byte, size)
|
||||||
|
for i := range buf {
|
||||||
|
buf[i] = byte(i)
|
||||||
|
}
|
||||||
|
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
||||||
|
b.SetBytes(int64(size))
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = Checksum(buf, 0)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
||||||
|
b.SetBytes(int64(size))
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = gvisorchecksum.Checksum(buf, 0)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that
|
||||||
|
// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical
|
||||||
|
// USO size (8192), and the kernel's max GSO superpacket (65535).
|
||||||
|
func BenchmarkChecksum(b *testing.B) {
|
||||||
|
for _, size := range []int{60, 1448, 8192, 65535} {
|
||||||
|
buf := make([]byte, size)
|
||||||
|
for i := range buf {
|
||||||
|
buf[i] = byte(i)
|
||||||
|
}
|
||||||
|
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
||||||
|
b.SetBytes(int64(size))
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = Checksum(buf, 0)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
||||||
|
b.SetBytes(int64(size))
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
_ = gvisorchecksum.Checksum(buf, 0)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+5
-9
@@ -8,8 +8,8 @@ import (
|
|||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
// defaultBatchBufSize is the per-Queue scratch size for Read. 65535 covers
|
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||||
// any single IP packet.
|
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||||
const defaultBatchBufSize = 65535
|
const defaultBatchBufSize = 65535
|
||||||
|
|
||||||
type Device interface {
|
type Device interface {
|
||||||
@@ -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)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
type offloadQueueSet struct {
|
||||||
|
pq []*Offload
|
||||||
|
// pqi is exactly the same as pq, but stored as the interface type
|
||||||
|
pqi []Queue
|
||||||
|
shutdownFd int
|
||||||
|
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6
|
||||||
|
// with the kernel. Queues created by Add inherit this and surface it
|
||||||
|
// via Offload.USOSupported so coalescers can gate USO emission.
|
||||||
|
usoEnabled bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
|
||||||
|
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
|
||||||
|
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
|
||||||
|
// should fall back to per-packet writes when this is false.
|
||||||
|
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
||||||
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &offloadQueueSet{
|
||||||
|
pq: []*Offload{},
|
||||||
|
pqi: []Queue{},
|
||||||
|
shutdownFd: shutdownFd,
|
||||||
|
usoEnabled: usoEnabled,
|
||||||
|
}
|
||||||
|
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *offloadQueueSet) Queues() []Queue {
|
||||||
|
return c.pqi
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *offloadQueueSet) Add(fd int) error {
|
||||||
|
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c.pq = append(c.pq, x)
|
||||||
|
c.pqi = append(c.pqi, x)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *offloadQueueSet) wakeForShutdown() error {
|
||||||
|
var buf [8]byte
|
||||||
|
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||||
|
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *offloadQueueSet) Close() error {
|
||||||
|
errs := []error{}
|
||||||
|
|
||||||
|
// Signal all readers blocked in poll to wake up and exit
|
||||||
|
if err := c.wakeForShutdown(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, x := range c.pq {
|
||||||
|
if err := x.Close(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -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...)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,65 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
|
||||||
|
// of pointer-capturing closure that sendInsideMessage builds.
|
||||||
|
type fakeBatch struct{ buf [65536]byte }
|
||||||
|
|
||||||
|
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
|
||||||
|
func (b *fakeBatch) Commit([]byte) {}
|
||||||
|
|
||||||
|
type fakeHostInfo struct {
|
||||||
|
remoteIndexId uint32
|
||||||
|
counter uint64
|
||||||
|
}
|
||||||
|
type fakeIface struct {
|
||||||
|
rebindCount uint8
|
||||||
|
hi *fakeHostInfo
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
|
||||||
|
// SegmentSuperpacket call when a closure captures pointer-bearing
|
||||||
|
// receivers — the realistic shape of sendInsideMessage's closure.
|
||||||
|
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
|
||||||
|
const mss = 1400
|
||||||
|
const numSeg = 32
|
||||||
|
pkt := buildTSOv6(mss*numSeg, mss)
|
||||||
|
gso := GSOInfo{
|
||||||
|
Size: mss,
|
||||||
|
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
|
||||||
|
CsumStart: 40,
|
||||||
|
Proto: GSOProtoTCP,
|
||||||
|
}
|
||||||
|
p := Packet{Bytes: pkt, GSO: gso}
|
||||||
|
|
||||||
|
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
|
||||||
|
f := &fakeIface{rebindCount: 7, hi: hi}
|
||||||
|
fb := &fakeBatch{}
|
||||||
|
|
||||||
|
// SegmentSuperpacket consumes pkt destructively; refresh from a master
|
||||||
|
// copy each iter (matches the production pattern where every TUN read
|
||||||
|
// hands the segmenter a fresh kernel-supplied buffer).
|
||||||
|
master := append([]byte(nil), pkt...)
|
||||||
|
work := make([]byte, len(pkt))
|
||||||
|
p.Bytes = work
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
copy(work, master)
|
||||||
|
err := SegmentSuperpacket(p, func(seg []byte) error {
|
||||||
|
out := fb.Reserve(16 + len(seg) + 16)
|
||||||
|
out[0] = byte(f.rebindCount)
|
||||||
|
out[1] = byte(hi.counter)
|
||||||
|
hi.counter++
|
||||||
|
fb.Commit(out)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
b.Fatalf("SegmentSuperpacket: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,22 @@
|
|||||||
|
//go:build !linux || android || e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import "fmt"
|
||||||
|
|
||||||
|
func protoFromGSOType(_ uint8) (GSOProto, error) {
|
||||||
|
return 0, fmt.Errorf("GSO unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
|
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
|
||||||
|
// builds (and Android/e2e_testing) this package does not provide a Queue
|
||||||
|
// implementation, so any caller that does construct a Packet here can only
|
||||||
|
// be operating on non-superpacket bytes and the stub forwards them
|
||||||
|
// directly. A non-zero GSO field is a programming error from the caller
|
||||||
|
// and returns an explicit error rather than silently misbehaving.
|
||||||
|
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
||||||
|
if pkt.GSO.IsSuperpacket() {
|
||||||
|
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
||||||
|
}
|
||||||
|
return fn(pkt.Bytes)
|
||||||
|
}
|
||||||
@@ -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()
|
|
||||||
}
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user