mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 13:06:57 +02:00
Compare commits
53 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f5ddff5ca1 | |||
| 400cbc26a1 | |||
| 01b31360df | |||
| 5bdf645b0b | |||
| 0375aff451 | |||
| 6cb00c613c | |||
| 40b4ae7fb4 | |||
| cf51b6dfd7 | |||
| fe93ebd017 | |||
| 961ddbfbc1 | |||
| 67bd9e848a | |||
| bc3f5d0400 | |||
| aef8e39cc4 | |||
| 69863d6c81 | |||
| 5d35351437 | |||
| f95857b4c3 | |||
| 398d67e2da | |||
| 696903d6d9 | |||
| c82db210ef | |||
| 1ada3d4dd9 | |||
| 5f920fdd7d | |||
| cba9ea5b1f | |||
| 83809a599a | |||
| 23c67bd8d8 | |||
| dd3a7ad03c | |||
| dd2ac5d655 | |||
| 76e82a5256 | |||
| eaf756ea6c | |||
| a82a8dc547 | |||
| 213dd46588 | |||
| 4fb5cdb4fa | |||
| ff91c37529 | |||
| b7e9939e92 | |||
| 33c2d7277c | |||
| f141cebe8d | |||
| 9ec8cf10f3 | |||
| 1ab1f71dba | |||
| d0f02ba873 | |||
| 5f890dbc34 | |||
| db85d61c23 | |||
| db9218b0be | |||
| 5f00ab4b74 | |||
| 2a1cc62001 | |||
| e753e6e93c | |||
| 32a7c04498 | |||
| 8c50fc3f60 | |||
| 2f4532f102 | |||
| 8c71f2f3f9 | |||
| 3d34cc9b74 | |||
| e80b9830a3 | |||
| 49e3c4649b | |||
| 72c04b90bd | |||
| 36ab1dbb97 |
@@ -0,0 +1,113 @@
|
||||
name: Code-sign Windows binaries
|
||||
description: >
|
||||
Sign every .exe under a given path in place via the DefinedNet code-signer
|
||||
Lambda. If `role` or `bucket` is empty, logs a notice and skips signing so
|
||||
forks and dev branches without AWS access still produce usable builds.
|
||||
|
||||
inputs:
|
||||
path:
|
||||
description: "Directory whose .exe files should be signed in place"
|
||||
required: true
|
||||
role:
|
||||
description: "IAM role ARN to assume via OIDC; empty disables signing"
|
||||
required: false
|
||||
default: ""
|
||||
bucket:
|
||||
description: "S3 staging bucket the code-signer Lambda reads from; empty disables signing"
|
||||
required: false
|
||||
default: ""
|
||||
region:
|
||||
description: "AWS region for the role and Lambda"
|
||||
required: false
|
||||
default: "us-east-2"
|
||||
function-name:
|
||||
description: "Code-signer Lambda function name"
|
||||
required: false
|
||||
default: "code-signer"
|
||||
key-prefix:
|
||||
description: "S3 key prefix the caller is authorized to write under"
|
||||
required: false
|
||||
default: "code-signing/slackhq/nebula"
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
steps:
|
||||
- name: Skip notice
|
||||
if: inputs.role == '' || inputs.bucket == ''
|
||||
shell: sh
|
||||
run: echo "::notice::code-signer role or bucket not set; skipping code signing."
|
||||
|
||||
- name: Configure AWS credentials
|
||||
if: inputs.role != '' && inputs.bucket != ''
|
||||
uses: aws-actions/configure-aws-credentials@v6
|
||||
with:
|
||||
role-to-assume: ${{ inputs.role }}
|
||||
aws-region: ${{ inputs.region }}
|
||||
# Default is 12 retries to ride out IAM trust-policy propagation; once
|
||||
# the role is stable we want a real misconfiguration to fail fast.
|
||||
retry-max-attempts: 5
|
||||
|
||||
- name: Sign .exe files
|
||||
if: inputs.role != '' && inputs.bucket != ''
|
||||
shell: sh
|
||||
env:
|
||||
SIGN_PATH: ${{ inputs.path }}
|
||||
BUCKET: ${{ inputs.bucket }}
|
||||
FUNCTION_NAME: ${{ inputs.function-name }}
|
||||
KEY_PREFIX: ${{ inputs.key-prefix }}
|
||||
run: |
|
||||
set -eu
|
||||
RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}"
|
||||
|
||||
find "$SIGN_PATH" -name '*.exe' -print | while read -r path
|
||||
do
|
||||
rel=${path#"$SIGN_PATH"/}
|
||||
file=$(basename "$path")
|
||||
name=${file%.exe}
|
||||
prefix="${KEY_PREFIX}/${RUN}"
|
||||
src="${prefix}/unsigned/${rel}"
|
||||
dst="${prefix}/signed/${rel}"
|
||||
|
||||
echo "::group::Sign ${rel}"
|
||||
echo "Uploading unsigned to s3://${BUCKET}/${src}"
|
||||
aws s3 cp --no-progress "$path" "s3://${BUCKET}/${src}" >/dev/null
|
||||
|
||||
echo "Invoking ${FUNCTION_NAME} Lambda"
|
||||
payload=$(jq -nc \
|
||||
--arg s "$src" \
|
||||
--arg d "$dst" \
|
||||
--arg p "$name" \
|
||||
'{source_key: $s, dest_key: $d, program_name: $p}')
|
||||
meta=$(aws lambda invoke \
|
||||
--function-name "$FUNCTION_NAME" \
|
||||
--cli-binary-format raw-in-base64-out \
|
||||
--payload "$payload" \
|
||||
--output json \
|
||||
/tmp/sign-resp.json)
|
||||
if echo "$meta" | jq -e '.FunctionError != null' >/dev/null
|
||||
then
|
||||
echo "::endgroup::"
|
||||
echo "::error::code-signer Lambda failed for ${rel}"
|
||||
cat /tmp/sign-resp.json >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Downloading signed back to ${path}"
|
||||
aws s3 cp --no-progress "s3://${BUCKET}/${dst}" "$path" >/dev/null
|
||||
|
||||
aws s3 rm "s3://${BUCKET}/${src}" >/dev/null 2>&1 || true
|
||||
aws s3 rm "s3://${BUCKET}/${dst}" >/dev/null 2>&1 || true
|
||||
|
||||
# Sanity-check the bytes we got back actually carry an Authenticode
|
||||
# signature that this machine can validate end to end.
|
||||
status=$(powershell -NoProfile -Command "(Get-AuthenticodeSignature -FilePath '$path').Status" | tr -d '\r')
|
||||
if [ "$status" != "Valid" ]
|
||||
then
|
||||
echo "::endgroup::"
|
||||
echo "::error::${rel} signature status: ${status} (expected Valid)"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "Signed ${rel} (sha256=$(jq -r '.sha256' /tmp/sign-resp.json), status=${status})"
|
||||
echo "::endgroup::"
|
||||
done
|
||||
@@ -24,7 +24,7 @@ jobs:
|
||||
mv build/*.tar.gz release
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v6
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: linux-latest
|
||||
path: release
|
||||
@@ -32,6 +32,9 @@ jobs:
|
||||
build-windows:
|
||||
name: Build Windows
|
||||
runs-on: windows-latest
|
||||
permissions:
|
||||
id-token: write
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
@@ -54,8 +57,15 @@ jobs:
|
||||
mkdir build\dist\windows
|
||||
mv dist\windows\wintun build\dist\windows\
|
||||
|
||||
- name: Code-sign
|
||||
uses: ./.github/actions/code-sign
|
||||
with:
|
||||
path: build
|
||||
role: ${{ secrets.DEFINED_CODE_SIGNER_ROLE }}
|
||||
bucket: ${{ secrets.DEFINED_CODE_SIGNER_BUCKET }}
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v6
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: windows-latest
|
||||
path: build
|
||||
@@ -75,7 +85,7 @@ jobs:
|
||||
|
||||
- name: Import certificates
|
||||
if: env.HAS_SIGNING_CREDS == 'true'
|
||||
uses: Apple-Actions/import-codesign-certs@v6
|
||||
uses: Apple-Actions/import-codesign-certs@v7
|
||||
with:
|
||||
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
||||
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
||||
@@ -104,7 +114,7 @@ jobs:
|
||||
fi
|
||||
|
||||
- name: Upload artifacts
|
||||
uses: actions/upload-artifact@v6
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: darwin-latest
|
||||
path: ./release/*
|
||||
@@ -128,21 +138,21 @@ jobs:
|
||||
|
||||
- name: Download artifacts
|
||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||
uses: actions/download-artifact@v7
|
||||
uses: actions/download-artifact@v8
|
||||
with:
|
||||
name: linux-latest
|
||||
path: artifacts
|
||||
|
||||
- name: Login to Docker Hub
|
||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||
uses: docker/login-action@v3
|
||||
uses: docker/login-action@v4
|
||||
with:
|
||||
username: ${{ vars.DOCKERHUB_USERNAME }}
|
||||
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||
uses: docker/setup-buildx-action@v3
|
||||
uses: docker/setup-buildx-action@v4
|
||||
|
||||
- name: Build and push images
|
||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||
@@ -163,7 +173,7 @@ jobs:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- name: Download artifacts
|
||||
uses: actions/download-artifact@v7
|
||||
uses: actions/download-artifact@v8
|
||||
with:
|
||||
path: artifacts
|
||||
|
||||
|
||||
@@ -14,10 +14,20 @@ on:
|
||||
- 'go.sum'
|
||||
jobs:
|
||||
|
||||
smoke-extra:
|
||||
smoke-extra-libvirt:
|
||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||
name: Run extra smoke tests
|
||||
name: ${{ matrix.target }}
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
target:
|
||||
- freebsd-amd64
|
||||
- openbsd-amd64
|
||||
- netbsd-amd64
|
||||
- linux-amd64-ipv6disable
|
||||
env:
|
||||
VAGRANT_DEFAULT_PROVIDER: libvirt
|
||||
steps:
|
||||
|
||||
- uses: actions/checkout@v6
|
||||
@@ -30,25 +40,93 @@ jobs:
|
||||
- name: add hashicorp source
|
||||
run: wget -O- https://apt.releases.hashicorp.com/gpg | gpg --dearmor | sudo tee /usr/share/keyrings/hashicorp-archive-keyring.gpg && echo "deb [signed-by=/usr/share/keyrings/hashicorp-archive-keyring.gpg] https://apt.releases.hashicorp.com $(lsb_release -cs) main" | sudo tee /etc/apt/sources.list.d/hashicorp.list
|
||||
|
||||
- name: workaround AMD-V issue # https://github.com/cri-o/packaging/pull/306
|
||||
run: sudo rmmod kvm_amd
|
||||
- name: install vagrant and libvirt
|
||||
run: |
|
||||
sudo apt-get update && sudo apt-get install -y vagrant libvirt-daemon-system libvirt-dev
|
||||
sudo chmod 666 /dev/kvm
|
||||
sudo usermod -aG libvirt $(whoami)
|
||||
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
||||
vagrant plugin install vagrant-libvirt
|
||||
|
||||
- name: install vagrant
|
||||
run: sudo apt-get update && sudo apt-get install -y vagrant virtualbox
|
||||
- name: ${{ matrix.target }}
|
||||
run: make smoke-vagrant/${{ matrix.target }}
|
||||
|
||||
- name: freebsd-amd64
|
||||
run: make smoke-vagrant/freebsd-amd64
|
||||
timeout-minutes: 30
|
||||
|
||||
- name: openbsd-amd64
|
||||
run: make smoke-vagrant/openbsd-amd64
|
||||
# linux-386 needs VirtualBox, which conflicts with KVM/libvirt -- isolated job.
|
||||
smoke-extra-virtualbox:
|
||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||
name: linux-386
|
||||
runs-on: ubuntu-latest
|
||||
env:
|
||||
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
||||
steps:
|
||||
|
||||
- name: netbsd-amd64
|
||||
run: make smoke-vagrant/netbsd-amd64
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
- name: add hashicorp source
|
||||
run: wget -O- https://apt.releases.hashicorp.com/gpg | gpg --dearmor | sudo tee /usr/share/keyrings/hashicorp-archive-keyring.gpg && echo "deb [signed-by=/usr/share/keyrings/hashicorp-archive-keyring.gpg] https://apt.releases.hashicorp.com $(lsb_release -cs) main" | sudo tee /etc/apt/sources.list.d/hashicorp.list
|
||||
|
||||
- name: install vagrant and virtualbox
|
||||
run: |
|
||||
sudo apt-get update && sudo apt-get install -y vagrant virtualbox
|
||||
sudo rmmod kvm_amd kvm_intel kvm 2>/dev/null || true
|
||||
|
||||
- name: linux-386
|
||||
run: make smoke-vagrant/linux-386
|
||||
|
||||
- name: linux-amd64-ipv6disable
|
||||
run: make smoke-vagrant/linux-amd64-ipv6disable
|
||||
|
||||
timeout-minutes: 30
|
||||
|
||||
smoke-windows:
|
||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||
name: Run windows smoke test
|
||||
runs-on: windows-latest
|
||||
steps:
|
||||
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
- uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version: '1.25'
|
||||
check-latest: true
|
||||
|
||||
# WSL2 + Ubuntu so the smoke can run a real linux peer with its own
|
||||
# netns. iputils-ping is needed for the in-WSL ping check. WSL1 has no
|
||||
# real kernel and would lack /dev/net/tun, so we have to force WSL2.
|
||||
- uses: Vampire/setup-wsl@v3
|
||||
with:
|
||||
distribution: Ubuntu-24.04
|
||||
additional-packages: iputils-ping iproute2
|
||||
|
||||
# Vampire/setup-wsl provisions WSL1 even when the WSL2 platform is present.
|
||||
# Convert the distro to WSL2 explicitly before we try to use /dev/net/tun.
|
||||
- name: convert distro to WSL2
|
||||
shell: pwsh
|
||||
run: |
|
||||
wsl --set-version Ubuntu-24.04 2
|
||||
wsl --shutdown
|
||||
wsl --list --verbose
|
||||
|
||||
- name: build windows nebula
|
||||
run: make bin-windows
|
||||
|
||||
- name: build linux nebula for WSL
|
||||
shell: bash
|
||||
env:
|
||||
GOOS: linux
|
||||
GOARCH: amd64
|
||||
run: |
|
||||
mkdir -p build/linux-amd64
|
||||
go build -o build/linux-amd64/nebula ./cmd/nebula
|
||||
|
||||
- name: run smoke-windows
|
||||
shell: pwsh
|
||||
working-directory: ./.github/workflows/smoke
|
||||
run: ./smoke-windows.ps1
|
||||
|
||||
timeout-minutes: 15
|
||||
|
||||
@@ -16,8 +16,10 @@ relay:
|
||||
am_relay: true
|
||||
EOF
|
||||
|
||||
export LIGHTHOUSES="192.168.100.1 172.17.0.2:4242"
|
||||
export REMOTE_ALLOW_LIST='{"172.17.0.4/32": false, "172.17.0.5/32": false}'
|
||||
# TEST-NET-3 placeholder IPs; smoke-relay.sh seds them to real container IPs.
|
||||
# Mapping: .2 lighthouse1, .3 host2, .4 host3, .5 host4.
|
||||
export LIGHTHOUSES="192.168.100.1 203.0.113.2:4242"
|
||||
export REMOTE_ALLOW_LIST='{"203.0.113.4/32": false, "203.0.113.5/32": false}'
|
||||
|
||||
HOST="host2" ../genconfig.sh >host2.yml <<EOF
|
||||
relay:
|
||||
@@ -25,7 +27,7 @@ relay:
|
||||
- 192.168.100.1
|
||||
EOF
|
||||
|
||||
export REMOTE_ALLOW_LIST='{"172.17.0.3/32": false}'
|
||||
export REMOTE_ALLOW_LIST='{"203.0.113.3/32": false}'
|
||||
|
||||
HOST="host3" ../genconfig.sh >host3.yml
|
||||
|
||||
|
||||
@@ -5,9 +5,15 @@ set -e -x
|
||||
rm -rf ./build
|
||||
mkdir ./build
|
||||
|
||||
# TODO: Assumes your docker bridge network is a /24, and the first container that launches will be .1
|
||||
# - We could make this better by launching the lighthouse first and then fetching what IP it is.
|
||||
NET="$(docker network inspect bridge -f '{{ range .IPAM.Config }}{{ .Subnet }}{{ end }}' | cut -d. -f1-3)"
|
||||
# Smoke containers run on a dedicated docker network whose subnet is allocated
|
||||
# at smoke time, not known at build time. Configs are written with TEST-NET-3
|
||||
# placeholder IPs (RFC 5737) and smoke.sh / smoke-vagrant.sh / smoke-relay.sh
|
||||
# sed the real container IPs in before starting nebula.
|
||||
#
|
||||
# Placeholder mapping (last octet == fixed container slot):
|
||||
# 203.0.113.2 -> lighthouse1, 203.0.113.3 -> host2,
|
||||
# 203.0.113.4 -> host3, 203.0.113.5 -> host4.
|
||||
LIGHTHOUSE_IP="203.0.113.2"
|
||||
|
||||
(
|
||||
cd build
|
||||
@@ -25,16 +31,16 @@ NET="$(docker network inspect bridge -f '{{ range .IPAM.Config }}{{ .Subnet }}{{
|
||||
../genconfig.sh >lighthouse1.yml
|
||||
|
||||
HOST="host2" \
|
||||
LIGHTHOUSES="192.168.100.1 $NET.2:4242" \
|
||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||
../genconfig.sh >host2.yml
|
||||
|
||||
HOST="host3" \
|
||||
LIGHTHOUSES="192.168.100.1 $NET.2:4242" \
|
||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||
INBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||
../genconfig.sh >host3.yml
|
||||
|
||||
HOST="host4" \
|
||||
LIGHTHOUSES="192.168.100.1 $NET.2:4242" \
|
||||
LIGHTHOUSES="192.168.100.1 $LIGHTHOUSE_IP:4242" \
|
||||
OUTBOUND='[{"port": "any", "proto": "icmp", "group": "lighthouse"}]' \
|
||||
../genconfig.sh >host4.yml
|
||||
|
||||
|
||||
@@ -6,6 +6,8 @@ set -o pipefail
|
||||
|
||||
mkdir -p logs
|
||||
|
||||
NETWORK="nebula-smoke-relay"
|
||||
|
||||
cleanup() {
|
||||
echo
|
||||
echo " *** cleanup"
|
||||
@@ -16,22 +18,53 @@ cleanup() {
|
||||
then
|
||||
docker kill lighthouse1 host2 host3 host4
|
||||
fi
|
||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
||||
}
|
||||
|
||||
trap cleanup EXIT
|
||||
|
||||
docker run --name lighthouse1 --rm nebula:smoke-relay -config lighthouse1.yml -test
|
||||
docker run --name host2 --rm nebula:smoke-relay -config host2.yml -test
|
||||
docker run --name host3 --rm nebula:smoke-relay -config host3.yml -test
|
||||
docker run --name host4 --rm nebula:smoke-relay -config host4.yml -test
|
||||
# Create a dedicated smoke network with an explicit subnet (required for --ip
|
||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
||||
# fail the whole test — we only need one to be free.
|
||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
||||
for candidate in 172.30.0.0/24 172.31.0.0/24 10.98.0.0/24 10.99.0.0/24 192.168.230.0/24; do
|
||||
if docker network create --subnet "$candidate" "$NETWORK" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
done
|
||||
if ! docker network inspect "$NETWORK" >/dev/null 2>&1; then
|
||||
echo "failed to create $NETWORK: every candidate subnet is in use" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
docker run --name lighthouse1 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||
# Derive container IPs from the network's assigned subnet. Slots: .2 lighthouse1,
|
||||
# .3 host2, .4 host3, .5 host4 — matches the placeholders in build-relay.sh.
|
||||
SUBNET="$(docker network inspect -f '{{(index .IPAM.Config 0).Subnet}}' "$NETWORK")"
|
||||
PREFIX="${SUBNET%/*}"
|
||||
PREFIX="${PREFIX%.*}"
|
||||
LIGHTHOUSE_IP="$PREFIX.2"
|
||||
HOST2_IP="$PREFIX.3"
|
||||
HOST3_IP="$PREFIX.4"
|
||||
HOST4_IP="$PREFIX.5"
|
||||
|
||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
||||
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
||||
sed "s|203\.0\.113\.|$PREFIX.|g" "$f" >"$f.tmp"
|
||||
mv "$f.tmp" "$f"
|
||||
done
|
||||
|
||||
docker run --name lighthouse1 --rm nebula:smoke-relay -config lighthouse1.yml -test
|
||||
docker run --name host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" nebula:smoke-relay -config host2.yml -test
|
||||
docker run --name host3 --rm -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" nebula:smoke-relay -config host3.yml -test
|
||||
docker run --name host4 --rm -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" nebula:smoke-relay -config host4.yml -test
|
||||
|
||||
docker run --name lighthouse1 --network "$NETWORK" --ip "$LIGHTHOUSE_IP" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||
sleep 1
|
||||
docker run --name host2 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||
docker run --name host2 --network "$NETWORK" --ip "$HOST2_IP" -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||
sleep 1
|
||||
docker run --name host3 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||
docker run --name host3 --network "$NETWORK" --ip "$HOST3_IP" -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||
sleep 1
|
||||
docker run --name host4 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
||||
docker run --name host4 --network "$NETWORK" --ip "$HOST4_IP" -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm nebula:smoke-relay -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
||||
sleep 1
|
||||
|
||||
set +x
|
||||
@@ -76,7 +109,13 @@ docker exec host4 sh -c 'kill 1'
|
||||
docker exec host3 sh -c 'kill 1'
|
||||
docker exec host2 sh -c 'kill 1'
|
||||
docker exec lighthouse1 sh -c 'kill 1'
|
||||
sleep 5
|
||||
|
||||
# Wait up to 30s for all backgrounded jobs to exit rather than relying on a
|
||||
# fixed sleep.
|
||||
for _ in $(seq 1 30); do
|
||||
[ -z "$(jobs -r)" ] && break
|
||||
sleep 1
|
||||
done
|
||||
|
||||
if [ "$(jobs -r)" ]
|
||||
then
|
||||
|
||||
@@ -8,6 +8,8 @@ export VAGRANT_CWD="$PWD/vagrant-$1"
|
||||
|
||||
mkdir -p logs
|
||||
|
||||
NETWORK="nebula-smoke"
|
||||
|
||||
cleanup() {
|
||||
echo
|
||||
echo " *** cleanup"
|
||||
@@ -19,32 +21,51 @@ cleanup() {
|
||||
docker kill lighthouse1 host2
|
||||
fi
|
||||
vagrant destroy -f
|
||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
||||
}
|
||||
|
||||
trap cleanup EXIT
|
||||
|
||||
# Create a dedicated smoke network with an explicit subnet (required for --ip
|
||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
||||
# fail the whole test — we only need one to be free.
|
||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
||||
for candidate in 172.30.0.0/24 172.31.0.0/24 10.98.0.0/24 10.99.0.0/24 192.168.230.0/24; do
|
||||
if docker network create --subnet "$candidate" "$NETWORK" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
done
|
||||
if ! docker network inspect "$NETWORK" >/dev/null 2>&1; then
|
||||
echo "failed to create $NETWORK: every candidate subnet is in use" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Derive container IPs from the network's assigned subnet. Slots: .2 lighthouse1,
|
||||
# .3 host2 — matches the placeholders in build.sh.
|
||||
SUBNET="$(docker network inspect -f '{{(index .IPAM.Config 0).Subnet}}' "$NETWORK")"
|
||||
PREFIX="${SUBNET%/*}"
|
||||
PREFIX="${PREFIX%.*}"
|
||||
LIGHTHOUSE_IP="$PREFIX.2"
|
||||
HOST2_IP="$PREFIX.3"
|
||||
|
||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
||||
# This must happen before `vagrant up` rsyncs build/ into the VM for host3.
|
||||
for f in build/host2.yml build/host3.yml; do
|
||||
sed "s|203\.0\.113\.|$PREFIX.|g" "$f" >"$f.tmp"
|
||||
mv "$f.tmp" "$f"
|
||||
done
|
||||
|
||||
CONTAINER="nebula:${NAME:-smoke}"
|
||||
|
||||
docker run --name lighthouse1 --rm "$CONTAINER" -config lighthouse1.yml -test
|
||||
docker run --name host2 --rm "$CONTAINER" -config host2.yml -test
|
||||
docker run --name host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" "$CONTAINER" -config host2.yml -test
|
||||
|
||||
vagrant up
|
||||
|
||||
# OpenBSD: synced folders are disabled because Vagrant's rsync installer
|
||||
# uses ftp.openbsd.org which no longer hosts packages for older releases.
|
||||
# Copy build artifacts in via scp instead.
|
||||
case "$1" in
|
||||
openbsd-*)
|
||||
vagrant ssh -c "sudo mkdir -p /nebula" -- -T
|
||||
tar -cf - -C build . | vagrant ssh -c "sudo tar -xf - -C /nebula && sudo chmod -R a+r /nebula" -- -T
|
||||
;;
|
||||
esac
|
||||
|
||||
vagrant ssh -c "cd /nebula && /nebula/$1-nebula -config host3.yml -test" -- -T
|
||||
|
||||
docker run --name lighthouse1 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||
docker run --name lighthouse1 --network "$NETWORK" --ip "$LIGHTHOUSE_IP" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||
sleep 1
|
||||
docker run --name host2 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||
docker run --name host2 --network "$NETWORK" --ip "$HOST2_IP" -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||
sleep 1
|
||||
vagrant ssh -c "cd /nebula && sudo sh -c 'echo \$\$ >/nebula/pid && exec /nebula/$1-nebula -config host3.yml'" 2>&1 -- -T | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||
sleep 15
|
||||
@@ -107,7 +128,14 @@ vagrant ssh -c "ping -c1 192.168.100.2" -- -T
|
||||
vagrant ssh -c "sudo xargs kill </nebula/pid" -- -T
|
||||
docker exec host2 sh -c 'kill 1'
|
||||
docker exec lighthouse1 sh -c 'kill 1'
|
||||
sleep 1
|
||||
|
||||
# Wait up to 30s for all backgrounded jobs to exit. vagrant ssh in particular
|
||||
# takes a beat to tear down after nebula exits on the VM, so a fixed sleep is
|
||||
# racy.
|
||||
for _ in $(seq 1 30); do
|
||||
[ -z "$(jobs -r)" ] && break
|
||||
sleep 1
|
||||
done
|
||||
|
||||
if [ "$(jobs -r)" ]
|
||||
then
|
||||
|
||||
@@ -0,0 +1,272 @@
|
||||
#!/usr/bin/env pwsh
|
||||
# Windows smoke test for the nebula tun + UDP + NLM code paths.
|
||||
#
|
||||
# Topology:
|
||||
# - lighthouse runs natively on the Windows host (wintun + windows UDP)
|
||||
# - peer runs inside WSL2 (Linux build of nebula, /dev/net/tun)
|
||||
#
|
||||
# WSL2 gives us a real netns boundary so the loopback fast-path on Windows
|
||||
# does not short-circuit the overlay -- when WSL pings the lighthouse VPN IP,
|
||||
# Linux has no idea that IP is local to the Windows host, so the packet is
|
||||
# forced through nebula. Same in reverse.
|
||||
|
||||
$ErrorActionPreference = 'Stop'
|
||||
|
||||
# wsl.exe emits UTF-16 LE by default which PowerShell reads as bytes, mangling
|
||||
# every captured string. WSL_UTF8 makes wsl.exe emit UTF-8 instead.
|
||||
$env:WSL_UTF8 = '1'
|
||||
|
||||
$RepoRoot = Resolve-Path "$PSScriptRoot\..\..\.."
|
||||
$Nebula = Join-Path $RepoRoot 'nebula.exe'
|
||||
$NebulaCert = Join-Path $RepoRoot 'nebula-cert.exe'
|
||||
$NebulaLinux = Join-Path $RepoRoot 'build\linux-amd64\nebula'
|
||||
|
||||
if (-not (Test-Path $Nebula)) { throw "missing $Nebula; run 'make bin-windows' first" }
|
||||
if (-not (Test-Path $NebulaCert)) { throw "missing $NebulaCert; run 'make bin-windows' first" }
|
||||
if (-not (Test-Path $NebulaLinux)) { throw "missing $NebulaLinux; build the linux nebula first" }
|
||||
|
||||
# Matches the distro installed by Vampire/setup-wsl in smoke-extra.yml.
|
||||
$Distro = 'Ubuntu-24.04'
|
||||
$listed = (wsl --list --quiet 2>$null) -join "`n"
|
||||
if ($listed -notmatch [regex]::Escape($Distro)) {
|
||||
throw "WSL distro $Distro not registered. Got: $listed"
|
||||
}
|
||||
Write-Host "Using WSL distro: $Distro"
|
||||
|
||||
# Windows host as seen from inside WSL: WSL's default-route gateway. We extract
|
||||
# it with a regex rather than awk fields so PowerShell does not eat any '$N'
|
||||
# tokens, and tabs/double-spaces in `ip route` output do not confuse a cut.
|
||||
$ipCmd = 'ip route show default | grep -oE "([0-9]+\.){3}[0-9]+" | head -1'
|
||||
$WindowsIp = (wsl -d $Distro -- bash -c $ipCmd).Trim()
|
||||
if (-not $WindowsIp) { throw "could not determine Windows host IP from WSL" }
|
||||
Write-Host "Windows host IP from WSL: $WindowsIp"
|
||||
|
||||
$WorkDir = Join-Path $env:TEMP 'nebula-smoke-windows'
|
||||
if (Test-Path $WorkDir) { Remove-Item -Recurse -Force $WorkDir }
|
||||
New-Item -ItemType Directory -Path $WorkDir | Out-Null
|
||||
|
||||
$WslDir = '/tmp/nebula-smoke'
|
||||
wsl -d $Distro -- bash -c "rm -rf $WslDir && mkdir -p $WslDir" | Out-Null
|
||||
|
||||
$DevName = 'nebula-smoke'
|
||||
$Ip1 = '192.168.241.1'
|
||||
$Ip2 = '192.168.241.2'
|
||||
$Port = 4242
|
||||
|
||||
& $NebulaCert ca -name 'smoke-ca' -out-crt "$WorkDir\ca.crt" -out-key "$WorkDir\ca.key"
|
||||
if ($LASTEXITCODE -ne 0) { throw "nebula-cert ca failed (exit $LASTEXITCODE)" }
|
||||
|
||||
& $NebulaCert sign -name 'lighthouse' -networks "$Ip1/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\lighthouse.crt" -out-key "$WorkDir\lighthouse.key"
|
||||
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign lighthouse failed (exit $LASTEXITCODE)" }
|
||||
|
||||
& $NebulaCert sign -name 'peer' -networks "$Ip2/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\peer.crt" -out-key "$WorkDir\peer.key"
|
||||
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign peer failed (exit $LASTEXITCODE)" }
|
||||
|
||||
# Windows lighthouse config.
|
||||
@"
|
||||
pki:
|
||||
ca: $WorkDir\ca.crt
|
||||
cert: $WorkDir\lighthouse.crt
|
||||
key: $WorkDir\lighthouse.key
|
||||
static_host_map: {}
|
||||
lighthouse:
|
||||
am_lighthouse: true
|
||||
interval: 60
|
||||
hosts: []
|
||||
listen:
|
||||
host: 0.0.0.0
|
||||
port: $Port
|
||||
tun:
|
||||
disabled: false
|
||||
dev: $DevName
|
||||
drop_local_broadcast: false
|
||||
drop_multicast: false
|
||||
tx_queue: 500
|
||||
mtu: 1300
|
||||
network_category: private
|
||||
logging:
|
||||
level: info
|
||||
format: text
|
||||
firewall:
|
||||
outbound_action: drop
|
||||
inbound_action: drop
|
||||
conntrack:
|
||||
tcp_timeout: 12m
|
||||
udp_timeout: 3m
|
||||
default_timeout: 10m
|
||||
outbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
inbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
"@ | Out-File -FilePath "$WorkDir\lighthouse.yml" -Encoding utf8
|
||||
|
||||
# WSL peer config (paths are POSIX, deliberately).
|
||||
@"
|
||||
pki:
|
||||
ca: $WslDir/ca.crt
|
||||
cert: $WslDir/peer.crt
|
||||
key: $WslDir/peer.key
|
||||
static_host_map:
|
||||
"${Ip1}": ["${WindowsIp}:$Port"]
|
||||
lighthouse:
|
||||
am_lighthouse: false
|
||||
interval: 60
|
||||
hosts:
|
||||
- "${Ip1}"
|
||||
listen:
|
||||
host: 0.0.0.0
|
||||
port: 0
|
||||
tun:
|
||||
disabled: false
|
||||
dev: nebula1
|
||||
drop_local_broadcast: false
|
||||
drop_multicast: false
|
||||
tx_queue: 500
|
||||
mtu: 1300
|
||||
logging:
|
||||
level: info
|
||||
format: text
|
||||
firewall:
|
||||
outbound_action: drop
|
||||
inbound_action: drop
|
||||
conntrack:
|
||||
tcp_timeout: 12m
|
||||
udp_timeout: 3m
|
||||
default_timeout: 10m
|
||||
outbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
inbound:
|
||||
- port: any
|
||||
proto: any
|
||||
host: any
|
||||
"@ | Out-File -FilePath "$WorkDir\peer.yml" -Encoding utf8
|
||||
|
||||
# Stage WSL artifacts. Convert Windows paths to WSL paths ourselves rather than
|
||||
# calling `wslpath`, because PowerShell's argument-passing to external EXEs
|
||||
# strips backslashes from path arguments in ways that are hard to escape around.
|
||||
function ConvertTo-WslPath {
|
||||
param([string]$WindowsPath)
|
||||
if ($WindowsPath -notmatch '^([A-Za-z]):\\(.*)$') {
|
||||
throw "cannot convert path to WSL: $WindowsPath"
|
||||
}
|
||||
return "/mnt/$($matches[1].ToLower())/$($matches[2].Replace('\','/'))"
|
||||
}
|
||||
|
||||
$WslWorkDir = ConvertTo-WslPath $WorkDir
|
||||
$WslNebulaPath = ConvertTo-WslPath $NebulaLinux
|
||||
wsl -d $Distro -- bash -c "cp '$WslWorkDir/ca.crt' '$WslWorkDir/peer.crt' '$WslWorkDir/peer.key' '$WslWorkDir/peer.yml' $WslDir/ && cp '$WslNebulaPath' $WslDir/nebula && chmod +x $WslDir/nebula"
|
||||
|
||||
# Make sure WSL has tun support and /dev/net/tun is usable before starting
|
||||
# nebula. Diagnostics first so a fail here points at the real problem (e.g.
|
||||
# WSL1 distros do not have a real kernel and will not have tun).
|
||||
Write-Host '=== WSL diagnostic ==='
|
||||
wsl --version 2>&1 | Out-Host
|
||||
wsl --list --verbose 2>&1 | Out-Host
|
||||
wsl -d $Distro -u root -- uname -a | Out-Host
|
||||
wsl -d $Distro -u root -- bash -c "modprobe tun 2>&1 || true; mkdir -p /dev/net; [ -c /dev/net/tun ] || mknod /dev/net/tun c 10 200; chmod 600 /dev/net/tun; ls -l /dev/net/tun"
|
||||
if ($LASTEXITCODE -ne 0) { throw "failed to prepare /dev/net/tun in WSL (TUN support missing?)" }
|
||||
|
||||
# Deliberately no New-NetFirewallRule calls here -- nebula's windows_bypass_wdf
|
||||
# feature is supposed to install WFP permit filters that let inbound traffic
|
||||
# through Windows Defender Firewall on its own. If this smoke regresses, that
|
||||
# feature regressed.
|
||||
|
||||
$lhOut = Join-Path $WorkDir 'lighthouse.out.log'
|
||||
$lhErr = Join-Path $WorkDir 'lighthouse.err.log'
|
||||
$lhProc = Start-Process -FilePath $Nebula -ArgumentList @('-config', "$WorkDir\lighthouse.yml") `
|
||||
-PassThru -NoNewWindow `
|
||||
-RedirectStandardOutput $lhOut `
|
||||
-RedirectStandardError $lhErr
|
||||
|
||||
# Run nebula in WSL as root with no sudo + no shell wrapper. PowerShell's
|
||||
# Start-Process arg quoting mangles `bash -c "..."` strings that contain
|
||||
# spaces/redirections, so we skip bash entirely and let Start-Process do the
|
||||
# stdout/stderr capture itself.
|
||||
$peerOut = Join-Path $WorkDir 'peer.out.log'
|
||||
$peerErr = Join-Path $WorkDir 'peer.err.log'
|
||||
$peerProc = Start-Process -FilePath 'wsl' `
|
||||
-ArgumentList @('-d', $Distro, '-u', 'root', '--', "$WslDir/nebula", '-config', "$WslDir/peer.yml") `
|
||||
-PassThru -NoNewWindow `
|
||||
-RedirectStandardOutput $peerOut `
|
||||
-RedirectStandardError $peerErr
|
||||
|
||||
function Wait-Until {
|
||||
param([scriptblock]$Predicate, [int]$TimeoutSec, [string]$What)
|
||||
$deadline = (Get-Date).AddSeconds($TimeoutSec)
|
||||
while ((Get-Date) -lt $deadline) {
|
||||
if (& $Predicate) { return }
|
||||
Start-Sleep -Milliseconds 500
|
||||
}
|
||||
throw "timed out waiting for: $What"
|
||||
}
|
||||
|
||||
try {
|
||||
Wait-Until -TimeoutSec 30 -What "windows wintun adapter $DevName with NetworkCategory=Private" -Predicate {
|
||||
if ($lhProc.HasExited) { throw "lighthouse exited (code $($lhProc.ExitCode)) before tun was ready" }
|
||||
$p = Get-NetConnectionProfile -InterfaceAlias $DevName -ErrorAction SilentlyContinue
|
||||
$p -and ("$($p.NetworkCategory)" -ieq 'Private')
|
||||
}
|
||||
Write-Host "OK: $DevName NetworkCategory=Private"
|
||||
|
||||
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
|
||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before tun was ready" }
|
||||
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet $Ip2' && echo yes"
|
||||
("$r").Trim() -eq 'yes'
|
||||
}
|
||||
Write-Host "OK: WSL nebula1 has $Ip2"
|
||||
|
||||
Wait-Until -TimeoutSec 30 -What "ping from WSL peer to windows lighthouse ($Ip1)" -Predicate {
|
||||
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before ping succeeded" }
|
||||
$r = wsl -d $Distro -u root -- bash -c "ping -c1 -W1 $Ip1 >/dev/null 2>&1 && echo OK"
|
||||
("$r").Trim() -eq 'OK'
|
||||
}
|
||||
Write-Host "OK: WSL peer -> windows lighthouse"
|
||||
|
||||
Wait-Until -TimeoutSec 30 -What "ping from windows lighthouse to WSL peer ($Ip2)" -Predicate {
|
||||
$null = & ping.exe -n 1 -w 1000 $Ip2
|
||||
$LASTEXITCODE -eq 0
|
||||
}
|
||||
Write-Host "OK: windows lighthouse -> WSL peer"
|
||||
|
||||
Write-Host ''
|
||||
Write-Host 'All smoke checks passed.'
|
||||
}
|
||||
catch {
|
||||
Write-Host ''
|
||||
Write-Host '=== lighthouse stdout ==='
|
||||
Get-Content $lhOut -ErrorAction SilentlyContinue | Out-Host
|
||||
Write-Host '=== lighthouse stderr ==='
|
||||
Get-Content $lhErr -ErrorAction SilentlyContinue | Out-Host
|
||||
Write-Host '=== peer stdout ==='
|
||||
Get-Content $peerOut -ErrorAction SilentlyContinue | Out-Host
|
||||
Write-Host '=== peer stderr ==='
|
||||
Get-Content $peerErr -ErrorAction SilentlyContinue | Out-Host
|
||||
Write-Host '=== nebula WFP filters ==='
|
||||
# Dump nebula-installed filters so we can verify they got registered with
|
||||
# the conditions we expect.
|
||||
$wfpDump = Join-Path $WorkDir 'wfp.xml'
|
||||
netsh wfp show filters file=$wfpDump 2>&1 | Out-Null
|
||||
if (Test-Path $wfpDump) {
|
||||
Select-String -Path $wfpDump -Pattern 'Nebula' -Context 0,80 -ErrorAction SilentlyContinue | Out-Host
|
||||
}
|
||||
throw
|
||||
}
|
||||
finally {
|
||||
if (-not $lhProc.HasExited) {
|
||||
Stop-Process -Id $lhProc.Id -Force -ErrorAction SilentlyContinue
|
||||
$lhProc.WaitForExit(5000) | Out-Null
|
||||
}
|
||||
wsl -d $Distro -u root -- bash -c "pkill -f $WslDir/nebula 2>/dev/null; true" | Out-Null
|
||||
# pkill returns 1 when no match and wsl propagates that; the smoke is done
|
||||
# so we don't want it to leak into the script's exit code.
|
||||
$global:LASTEXITCODE = 0
|
||||
if ($peerProc -and -not $peerProc.HasExited) {
|
||||
Stop-Process -Id $peerProc.Id -Force -ErrorAction SilentlyContinue
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,8 @@ set -o pipefail
|
||||
|
||||
mkdir -p logs
|
||||
|
||||
NETWORK="nebula-smoke"
|
||||
|
||||
cleanup() {
|
||||
echo
|
||||
echo " *** cleanup"
|
||||
@@ -16,24 +18,56 @@ cleanup() {
|
||||
then
|
||||
docker kill lighthouse1 host2 host3 host4
|
||||
fi
|
||||
docker network rm "$NETWORK" >/dev/null 2>&1
|
||||
}
|
||||
|
||||
trap cleanup EXIT
|
||||
|
||||
# Create a dedicated smoke network with an explicit subnet (required for --ip
|
||||
# below). Probe a short list of candidates so a locally-used range doesn't
|
||||
# fail the whole test — we only need one to be free.
|
||||
docker network rm "$NETWORK" >/dev/null 2>&1 || true
|
||||
for candidate in 172.30.0.0/24 172.31.0.0/24 10.98.0.0/24 10.99.0.0/24 192.168.230.0/24; do
|
||||
if docker network create --subnet "$candidate" "$NETWORK" >/dev/null 2>&1; then
|
||||
break
|
||||
fi
|
||||
done
|
||||
if ! docker network inspect "$NETWORK" >/dev/null 2>&1; then
|
||||
echo "failed to create $NETWORK: every candidate subnet is in use" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Derive container IPs from the network's assigned subnet. Slots: .2 lighthouse1,
|
||||
# .3 host2, .4 host3, .5 host4 — matches the placeholders in build.sh.
|
||||
SUBNET="$(docker network inspect -f '{{(index .IPAM.Config 0).Subnet}}' "$NETWORK")"
|
||||
PREFIX="${SUBNET%/*}"
|
||||
PREFIX="${PREFIX%.*}"
|
||||
LIGHTHOUSE_IP="$PREFIX.2"
|
||||
HOST2_IP="$PREFIX.3"
|
||||
HOST3_IP="$PREFIX.4"
|
||||
HOST4_IP="$PREFIX.5"
|
||||
|
||||
# Sed the placeholder TEST-NET-3 IPs in the host configs to the real ones.
|
||||
# build/lighthouse1.yml has no IPs to rewrite so it's skipped.
|
||||
for f in build/host2.yml build/host3.yml build/host4.yml; do
|
||||
sed "s|203\.0\.113\.|$PREFIX.|g" "$f" >"$f.tmp"
|
||||
mv "$f.tmp" "$f"
|
||||
done
|
||||
|
||||
CONTAINER="nebula:${NAME:-smoke}"
|
||||
|
||||
docker run --name lighthouse1 --rm "$CONTAINER" -config lighthouse1.yml -test
|
||||
docker run --name host2 --rm "$CONTAINER" -config host2.yml -test
|
||||
docker run --name host3 --rm "$CONTAINER" -config host3.yml -test
|
||||
docker run --name host4 --rm "$CONTAINER" -config host4.yml -test
|
||||
docker run --name host2 --rm -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" "$CONTAINER" -config host2.yml -test
|
||||
docker run --name host3 --rm -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" "$CONTAINER" -config host3.yml -test
|
||||
docker run --name host4 --rm -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" "$CONTAINER" -config host4.yml -test
|
||||
|
||||
docker run --name lighthouse1 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||
docker run --name lighthouse1 --network "$NETWORK" --ip "$LIGHTHOUSE_IP" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config lighthouse1.yml 2>&1 | tee logs/lighthouse1 | sed -u 's/^/ [lighthouse1] /' &
|
||||
sleep 1
|
||||
docker run --name host2 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||
docker run --name host2 --network "$NETWORK" --ip "$HOST2_IP" -v "$PWD/build/host2.yml:/nebula/host2.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host2.yml 2>&1 | tee logs/host2 | sed -u 's/^/ [host2] /' &
|
||||
sleep 1
|
||||
docker run --name host3 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||
docker run --name host3 --network "$NETWORK" --ip "$HOST3_IP" -v "$PWD/build/host3.yml:/nebula/host3.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host3.yml 2>&1 | tee logs/host3 | sed -u 's/^/ [host3] /' &
|
||||
sleep 1
|
||||
docker run --name host4 --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
||||
docker run --name host4 --network "$NETWORK" --ip "$HOST4_IP" -v "$PWD/build/host4.yml:/nebula/host4.yml:ro" --device /dev/net/tun:/dev/net/tun --cap-add NET_ADMIN --rm "$CONTAINER" -config host4.yml 2>&1 | tee logs/host4 | sed -u 's/^/ [host4] /' &
|
||||
sleep 1
|
||||
|
||||
# grab tcpdump pcaps for debugging
|
||||
@@ -48,7 +82,7 @@ docker exec host4 tcpdump -i eth0 -q -w - -U 2>logs/host4.outside.log >logs/host
|
||||
|
||||
docker exec host2 ncat -nklv 0.0.0.0 2000 &
|
||||
docker exec host3 ncat -nklv 0.0.0.0 2000 &
|
||||
docker exec host4 ncat -nkluv 0.0.0.0 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 0.0.0.0 3000 &
|
||||
docker exec host3 ncat -e '/usr/bin/echo host3' -nkluv 0.0.0.0 3000 &
|
||||
|
||||
@@ -121,17 +155,23 @@ echo " *** Testing conntrack"
|
||||
echo
|
||||
set -x
|
||||
|
||||
# host2 speaking to host4 on UDP 4000 should allow it to reply, when firewall rules would normally not permit this
|
||||
docker exec host2 sh -c "/usr/bin/echo host2 | ncat -nuv 192.168.100.4 4000"
|
||||
docker exec host2 ncat -e '/usr/bin/echo helloagainfromhost2' -nkluv 0.0.0.0 4000 &
|
||||
sleep 1
|
||||
docker exec host4 sh -c "/usr/bin/echo host4 | ncat -nuv 192.168.100.2 4000"
|
||||
# host4's outbound firewall only allows ICMP to the lighthouse, so host4
|
||||
# cannot initiate UDP to host2. Once host2 initiates a flow to host4:4000,
|
||||
# conntrack must let host4's listener reply on that flow. If it doesn't,
|
||||
# the echo back from host4 never reaches host2.
|
||||
docker exec host2 sh -c "(/usr/bin/echo host2; sleep 2) | ncat -nuv 192.168.100.4 4000" | grep -q helloagainfromhost4
|
||||
|
||||
docker exec host4 sh -c 'kill 1'
|
||||
docker exec host3 sh -c 'kill 1'
|
||||
docker exec host2 sh -c 'kill 1'
|
||||
docker exec lighthouse1 sh -c 'kill 1'
|
||||
sleep 5
|
||||
|
||||
# Wait up to 30s for all backgrounded jobs to exit rather than relying on a
|
||||
# fixed sleep.
|
||||
for _ in $(seq 1 30); do
|
||||
[ -z "$(jobs -r)" ] && break
|
||||
sleep 1
|
||||
done
|
||||
|
||||
if [ "$(jobs -r)" ]
|
||||
then
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# -*- mode: ruby -*-
|
||||
# vi: set ft=ruby :
|
||||
Vagrant.configure("2") do |config|
|
||||
config.vm.box = "ubuntu/jammy64"
|
||||
config.vm.box = "bento/ubuntu-24.04"
|
||||
|
||||
config.vm.synced_folder "../build", "/nebula"
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# -*- mode: ruby -*-
|
||||
# vi: set ft=ruby :
|
||||
Vagrant.configure("2") do |config|
|
||||
config.vm.box = "generic/netbsd9"
|
||||
config.vm.box = "DefinedNet/netbsd10"
|
||||
|
||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
||||
end
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# -*- mode: ruby -*-
|
||||
# vi: set ft=ruby :
|
||||
Vagrant.configure("2") do |config|
|
||||
config.vm.box = "generic/openbsd7"
|
||||
config.vm.box = "DefinedNet/openbsd78"
|
||||
|
||||
config.vm.synced_folder ".", "/vagrant", disabled: true
|
||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
||||
end
|
||||
|
||||
@@ -45,7 +45,7 @@ jobs:
|
||||
- name: Build test mobile
|
||||
run: make build-test-mobile
|
||||
|
||||
- uses: actions/upload-artifact@v6
|
||||
- uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: e2e packet flow linux-latest
|
||||
path: e2e/mermaid/linux-latest
|
||||
@@ -125,7 +125,7 @@ jobs:
|
||||
- name: End 2 end
|
||||
run: make e2evv
|
||||
|
||||
- uses: actions/upload-artifact@v6
|
||||
- uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: e2e packet flow ${{ matrix.os }}
|
||||
path: e2e/mermaid/${{ matrix.os }}
|
||||
|
||||
@@ -2,7 +2,21 @@ version: "2"
|
||||
linters:
|
||||
default: none
|
||||
enable:
|
||||
- sloglint
|
||||
- testifylint
|
||||
settings:
|
||||
sloglint:
|
||||
# Enforce key-value pair form for Info/Debug/Warn/Error/Log/With and
|
||||
# the package-level slog equivalents. Use l.Log(ctx, level, ...) for
|
||||
# custom levels instead of LogAttrs when you can.
|
||||
#
|
||||
# LogAttrs is also flagged by this rule because it takes ...slog.Attr;
|
||||
# the few legitimate sites (where attrs is built up as a []slog.Attr)
|
||||
# carry a //nolint:sloglint with rationale.
|
||||
kv-only: true
|
||||
# no-mixed-args is on by default: forbids mixing kv and attrs in one call.
|
||||
# discard-handler is on by default (since Go 1.24): suggests
|
||||
# slog.DiscardHandler over slog.NewTextHandler(io.Discard, nil).
|
||||
exclusions:
|
||||
generated: lax
|
||||
presets:
|
||||
|
||||
@@ -1,23 +1,43 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math"
|
||||
mathbits "math/bits"
|
||||
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const bitsPerWord = 64
|
||||
|
||||
// Bits is a sliding-window anti-replay tracker. The window is stored as a
|
||||
// circular bitmap packed into uint64 words (8x denser than a []bool), so a
|
||||
// length-N window costs N/8 bytes. length must be a power of two.
|
||||
type Bits struct {
|
||||
length uint64
|
||||
lengthMask uint64
|
||||
current uint64
|
||||
bits []bool
|
||||
bits []uint64
|
||||
lostCounter metrics.Counter
|
||||
dupeCounter metrics.Counter
|
||||
outOfWindowCounter metrics.Counter
|
||||
}
|
||||
|
||||
func NewBits(bits uint64) *Bits {
|
||||
func NewBits(length uint64) *Bits {
|
||||
if length == 0 || length&(length-1) != 0 {
|
||||
panic(fmt.Sprintf("Bits length must be a power of two, got %d", length))
|
||||
}
|
||||
|
||||
nWords := length / bitsPerWord
|
||||
if nWords == 0 {
|
||||
nWords = 1
|
||||
}
|
||||
b := &Bits{
|
||||
length: bits,
|
||||
bits: make([]bool, bits, bits),
|
||||
length: length,
|
||||
lengthMask: length - 1,
|
||||
bits: make([]uint64, nWords),
|
||||
current: 0,
|
||||
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
||||
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
||||
@@ -25,88 +45,219 @@ func NewBits(bits uint64) *Bits {
|
||||
}
|
||||
|
||||
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
||||
b.bits[0] = true
|
||||
b.current = 0
|
||||
b.bits[0] = 1
|
||||
return b
|
||||
}
|
||||
|
||||
func (b *Bits) Check(l *logrus.Logger, i uint64) bool {
|
||||
func (b *Bits) get(i uint64) bool {
|
||||
pos := i & b.lengthMask
|
||||
//bit-shifting by 6 because i is a bit index, not a u64 index, and we need to find the u64 without bit in it
|
||||
return b.bits[pos>>6]&(uint64(1)<<(pos&63)) != 0
|
||||
}
|
||||
|
||||
func (b *Bits) set(i uint64) {
|
||||
pos := i & b.lengthMask
|
||||
b.bits[pos>>6] |= uint64(1) << (pos & 63)
|
||||
}
|
||||
|
||||
// clearRange clears `count` bits starting at circular position `startPos`
|
||||
// (already masked to [0, length)) and returns how many of them were set
|
||||
// before the clear. count must be in [1, length].
|
||||
func (b *Bits) clearRange(startPos, count uint64) uint64 {
|
||||
wasSet := uint64(0)
|
||||
if count >= b.length {
|
||||
for _, w := range b.bits {
|
||||
wasSet += uint64(mathbits.OnesCount64(w))
|
||||
}
|
||||
clear(b.bits)
|
||||
return wasSet
|
||||
}
|
||||
|
||||
pos := startPos
|
||||
remaining := count
|
||||
|
||||
// handle the potential partial word before pos becomes u64 aligned
|
||||
word := pos >> 6
|
||||
bit := pos & 63
|
||||
take := uint64(64) - bit
|
||||
if take > remaining {
|
||||
take = remaining
|
||||
}
|
||||
if take > b.length-pos {
|
||||
take = b.length - pos
|
||||
}
|
||||
var mask uint64
|
||||
if take == 64 {
|
||||
mask = math.MaxUint64
|
||||
} else {
|
||||
mask = ((uint64(1) << take) - 1) << bit
|
||||
}
|
||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word] & mask))
|
||||
b.bits[word] &^= mask
|
||||
remaining -= take
|
||||
pos = (pos + take) & b.lengthMask
|
||||
|
||||
// Clear whole words, keeping track of the number of set bits
|
||||
for remaining >= 64 {
|
||||
word = pos >> 6
|
||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word]))
|
||||
b.bits[word] = 0
|
||||
remaining -= 64
|
||||
pos = (pos + 64) & b.lengthMask
|
||||
}
|
||||
|
||||
// Clear the remaining partial word
|
||||
if remaining > 0 {
|
||||
word = pos >> 6
|
||||
mask = (uint64(1) << remaining) - 1
|
||||
wasSet += uint64(mathbits.OnesCount64(b.bits[word] & mask))
|
||||
b.bits[word] &^= mask
|
||||
}
|
||||
|
||||
return wasSet
|
||||
}
|
||||
|
||||
func (b *Bits) strictlyWithinWindow(i uint64) bool {
|
||||
// Handle the case where the window hasn't slid yet. This avoids u64 underflow.
|
||||
inWarmup := b.current < b.length
|
||||
if i < b.length && inWarmup {
|
||||
return true
|
||||
}
|
||||
|
||||
// Next, if the packet is in-window, see if we've seen it before
|
||||
if i > b.current-b.length {
|
||||
return true
|
||||
}
|
||||
return false //not within window!
|
||||
}
|
||||
|
||||
// Check returns true if i is within (or way out in front of) the window, and not a replay
|
||||
func (b *Bits) Check(l *slog.Logger, i uint64) bool {
|
||||
// If i is the next number, return true.
|
||||
if i > b.current {
|
||||
return true
|
||||
}
|
||||
|
||||
// If i is within the window, check if it's been set already.
|
||||
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||
return !b.bits[i%b.length]
|
||||
if b.strictlyWithinWindow(i) {
|
||||
return !b.get(i)
|
||||
}
|
||||
|
||||
// Not within the window
|
||||
if l.Level >= logrus.DebugLevel {
|
||||
l.Debugf("rejected a packet (top) %d %d\n", b.current, i)
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
l.Debug("rejected a packet (top)", "current", b.current, "incoming", i)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (b *Bits) Update(l *logrus.Logger, i uint64) bool {
|
||||
// If i is the next number, return true and update current.
|
||||
// Update has three branches:
|
||||
// - i == b.current+1: fast path; advance the cursor by one and lose-count
|
||||
// the slot we just stomped (only past warmup; see the i > b.length guard
|
||||
// below).
|
||||
// - i > b.current+1: jump path; clear all slots between current and i
|
||||
// (or up to a full window's worth, whichever is smaller) via clearRange,
|
||||
// then mark i. Two arms here: a warmup arm that handles the very first
|
||||
// window before the cursor has slid, and a steady-state arm that treats
|
||||
// every cleared empty slot as a lost packet.
|
||||
// - i <= b.current: in-window check for duplicates; out-of-window otherwise.
|
||||
//
|
||||
// NewBits seeds bits[0]=1 so counter 0 looks "received" — Update never
|
||||
// clears that marker during warmup (clearRange skips position 0 when
|
||||
// startPos=1), and once b.current >= b.length the marker is no longer
|
||||
// consulted. The marker prevents a fictitious "lost" hit on the first real
|
||||
// counter.
|
||||
func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
||||
// Fast path: i is the next expected counter. Split out so the function
|
||||
// stays small and avoids paying for the slow paths' slog argument-build
|
||||
// stack frame on every call. The bit read/test/write is inlined to
|
||||
// touch the backing word once.
|
||||
if i == b.current+1 {
|
||||
// Check if the oldest bit was lost since we are shifting the window by 1 and occupying it with this counter
|
||||
// The very first window can only be tracked as lost once we are on the 2nd window or greater
|
||||
if b.bits[i%b.length] == false && i > b.length {
|
||||
pos := i & b.lengthMask
|
||||
word := pos >> 6
|
||||
mask := uint64(1) << (pos & 63)
|
||||
w := b.bits[word]
|
||||
if i > b.length && w&mask == 0 {
|
||||
b.lostCounter.Inc(1)
|
||||
}
|
||||
b.bits[i%b.length] = true
|
||||
b.bits[word] = w | mask
|
||||
b.current = i
|
||||
return true
|
||||
}
|
||||
return b.updateSlow(l, i)
|
||||
}
|
||||
|
||||
// updateSlow handles jumps, in-window backfill, dupes, and out-of-window.
|
||||
func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
|
||||
// If i is a jump, adjust the window, record lost, update current, and return true
|
||||
if i > b.current {
|
||||
lost := int64(0)
|
||||
// Zero out the bits between the current and the new counter value, limited by the window size,
|
||||
// since the window is shifting
|
||||
for n := b.current + 1; n <= min(i, b.current+b.length); n++ {
|
||||
if b.bits[n%b.length] == false && n > b.length {
|
||||
lost++
|
||||
end := i
|
||||
if end > b.current+b.length {
|
||||
end = b.current + b.length
|
||||
}
|
||||
count := end - b.current
|
||||
startPos := (b.current + 1) & b.lengthMask
|
||||
|
||||
var lost int64
|
||||
if b.current >= b.length {
|
||||
// Steady state: every cleared slot is past warmup, so any unset
|
||||
// bit we evict is a lost packet from the previous cycle.
|
||||
wasSet := b.clearRange(startPos, count)
|
||||
lost = int64(count) - int64(wasSet)
|
||||
} else {
|
||||
// Warmup (the very first window). Some cleared slots represent
|
||||
// packets <= length where eviction is not "lost" in the usual
|
||||
// sense. This branch is taken at most once per connection so we
|
||||
// don't bother optimizing it.
|
||||
for n := b.current + 1; n <= end; n++ {
|
||||
if !b.get(n) && n > b.length {
|
||||
lost++
|
||||
}
|
||||
}
|
||||
b.bits[n%b.length] = false
|
||||
b.clearRange(startPos, count)
|
||||
}
|
||||
|
||||
// Only record any skipped packets as a result of the window moving further than the window length
|
||||
// Any loss within the new window will be accounted for in future calls
|
||||
lost += max(0, int64(i-b.current-b.length))
|
||||
// Anything past the new window can never be backfilled, so it's lost.
|
||||
if i > b.current+b.length {
|
||||
lost += int64(i - b.current - b.length)
|
||||
}
|
||||
b.lostCounter.Inc(lost)
|
||||
|
||||
b.bits[i%b.length] = true
|
||||
b.set(i)
|
||||
b.current = i
|
||||
return true
|
||||
}
|
||||
|
||||
// If i is within the current window but below the current counter,
|
||||
// Check to see if it's a duplicate
|
||||
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||
if b.current == i || b.bits[i%b.length] == true {
|
||||
if l.Level >= logrus.DebugLevel {
|
||||
l.WithField("receiveWindow", m{"accepted": false, "currentCounter": b.current, "incomingCounter": i, "reason": "duplicate"}).
|
||||
Debug("Receive window")
|
||||
// If i is within the current window but below the current counter, check to see if it's a duplicate
|
||||
if b.strictlyWithinWindow(i) {
|
||||
pos := i & b.lengthMask
|
||||
word := pos >> 6
|
||||
mask := uint64(1) << (pos & 63)
|
||||
w := b.bits[word]
|
||||
if b.current == i || w&mask != 0 {
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
l.Debug("Receive window",
|
||||
"accepted", false,
|
||||
"currentCounter", b.current,
|
||||
"incomingCounter", i,
|
||||
"reason", "duplicate",
|
||||
)
|
||||
}
|
||||
b.dupeCounter.Inc(1)
|
||||
return false
|
||||
}
|
||||
|
||||
b.bits[i%b.length] = true
|
||||
b.bits[word] = w | mask
|
||||
return true
|
||||
}
|
||||
|
||||
// In all other cases, fail and don't change current.
|
||||
b.outOfWindowCounter.Inc(1)
|
||||
if l.Level >= logrus.DebugLevel {
|
||||
l.WithField("accepted", false).
|
||||
WithField("currentCounter", b.current).
|
||||
WithField("incomingCounter", i).
|
||||
WithField("reason", "nonsense").
|
||||
Debug("Receive window")
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
l.Debug("Receive window",
|
||||
"accepted", false,
|
||||
"currentCounter", b.current,
|
||||
"incomingCounter", i,
|
||||
"reason", "nonsense",
|
||||
)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
+277
-130
@@ -7,61 +7,79 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
// snapshot returns the bitmap as a []bool of length b.length, for readable
|
||||
// test assertions against the now-packed []uint64 storage.
|
||||
func (b *Bits) snapshot() []bool {
|
||||
out := make([]bool, b.length)
|
||||
for i := uint64(0); i < b.length; i++ {
|
||||
out[i] = b.get(i)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func TestBitsRequiresPowerOfTwo(t *testing.T) {
|
||||
assert.Panics(t, func() { NewBits(10) })
|
||||
assert.Panics(t, func() { NewBits(0) })
|
||||
assert.NotPanics(t, func() { NewBits(1) })
|
||||
assert.NotPanics(t, func() { NewBits(16) })
|
||||
assert.NotPanics(t, func() { NewBits(1024) })
|
||||
assert.NotPanics(t, func() { NewBits(16384) })
|
||||
}
|
||||
|
||||
func TestBits(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(10)
|
||||
|
||||
// make sure it is the right size
|
||||
assert.Len(t, b.bits, 10)
|
||||
b := NewBits(16)
|
||||
assert.EqualValues(t, 16, b.length)
|
||||
|
||||
// This is initialized to zero - receive one. This should work.
|
||||
assert.True(t, b.Check(l, 1))
|
||||
assert.True(t, b.Update(l, 1))
|
||||
assert.EqualValues(t, 1, b.current)
|
||||
g := []bool{true, true, false, false, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.bits)
|
||||
g := []bool{true, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.snapshot())
|
||||
|
||||
// Receive two
|
||||
assert.True(t, b.Check(l, 2))
|
||||
assert.True(t, b.Update(l, 2))
|
||||
assert.EqualValues(t, 2, b.current)
|
||||
g = []bool{true, true, true, false, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.bits)
|
||||
g = []bool{true, true, true, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.snapshot())
|
||||
|
||||
// Receive two again - it will fail
|
||||
assert.False(t, b.Check(l, 2))
|
||||
assert.False(t, b.Update(l, 2))
|
||||
assert.EqualValues(t, 2, b.current)
|
||||
|
||||
// Jump ahead to 15, which should clear everything and set the 6th element
|
||||
assert.True(t, b.Check(l, 15))
|
||||
assert.True(t, b.Update(l, 15))
|
||||
assert.EqualValues(t, 15, b.current)
|
||||
g = []bool{false, false, false, false, false, true, false, false, false, false}
|
||||
assert.Equal(t, g, b.bits)
|
||||
// Jump ahead to 25, which clears the window and sets slot 25%16 = 9.
|
||||
assert.True(t, b.Check(l, 25))
|
||||
assert.True(t, b.Update(l, 25))
|
||||
assert.EqualValues(t, 25, b.current)
|
||||
g = []bool{false, false, false, false, false, false, false, false, false, true, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.snapshot())
|
||||
|
||||
// Mark 14, which is allowed because it is in the window
|
||||
assert.True(t, b.Check(l, 14))
|
||||
assert.True(t, b.Update(l, 14))
|
||||
assert.EqualValues(t, 15, b.current)
|
||||
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||
assert.Equal(t, g, b.bits)
|
||||
// Mark 24, which is in window (current 25, length 16, window covers [10,25]).
|
||||
assert.True(t, b.Check(l, 24))
|
||||
assert.True(t, b.Update(l, 24))
|
||||
assert.EqualValues(t, 25, b.current)
|
||||
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.snapshot())
|
||||
|
||||
// Mark 5, which is not allowed because it is not in the window
|
||||
// Mark 5, not allowed because 5 <= current-length (25-16=9).
|
||||
assert.False(t, b.Check(l, 5))
|
||||
assert.False(t, b.Update(l, 5))
|
||||
assert.EqualValues(t, 15, b.current)
|
||||
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||
assert.Equal(t, g, b.bits)
|
||||
assert.EqualValues(t, 25, b.current)
|
||||
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
||||
assert.Equal(t, g, b.snapshot())
|
||||
|
||||
// make sure we handle wrapping around once to the current position
|
||||
b = NewBits(10)
|
||||
// Make sure we handle wrapping around once to the same slot. With
|
||||
// length=16, packets 1 and 17 share slot 1.
|
||||
b = NewBits(16)
|
||||
assert.True(t, b.Update(l, 1))
|
||||
assert.True(t, b.Update(l, 11))
|
||||
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false}, b.bits)
|
||||
assert.True(t, b.Update(l, 17))
|
||||
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}, b.snapshot())
|
||||
|
||||
// Walk through a few windows in order
|
||||
b = NewBits(10)
|
||||
b = NewBits(16)
|
||||
for i := uint64(1); i <= 100; i++ {
|
||||
assert.True(t, b.Check(l, i), "Error while checking %v", i)
|
||||
assert.True(t, b.Update(l, i), "Error while updating %v", i)
|
||||
@@ -72,24 +90,31 @@ func TestBits(t *testing.T) {
|
||||
|
||||
func TestBitsLargeJumps(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(10)
|
||||
|
||||
// length=16. Update(55) from current=0:
|
||||
// warmup, per-bit loop sees no n>16 with unset bits (slot 0 was set by
|
||||
// NewBits and gets re-evaluated when n=16; n=16 is not strictly > 16),
|
||||
// so the loop contributes 0. The jump exceeds the window so we record
|
||||
// 55 - 0 - 16 = 39 packets fell out the back.
|
||||
b := NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
assert.True(t, b.Update(l, 55))
|
||||
assert.Equal(t, int64(39), b.lostCounter.Count())
|
||||
|
||||
b = NewBits(10)
|
||||
b.lostCounter.Clear()
|
||||
assert.True(t, b.Update(l, 55)) // We saw packet 55 and can still track 45,46,47,48,49,50,51,52,53,54
|
||||
assert.Equal(t, int64(45), b.lostCounter.Count())
|
||||
// Update(100): clears 16 slots starting at slot 56%16=8. Only slot 7 (for
|
||||
// packet 55) was set, so 16 - 1 = 15 evicted slots had unset bits.
|
||||
// Plus 100 - 55 - 16 = 29 packets fell past the window. Total 44.
|
||||
assert.True(t, b.Update(l, 100))
|
||||
assert.Equal(t, int64(39+44), b.lostCounter.Count())
|
||||
|
||||
assert.True(t, b.Update(l, 100)) // We saw packet 55 and 100 and can still track 90,91,92,93,94,95,96,97,98,99
|
||||
assert.Equal(t, int64(89), b.lostCounter.Count())
|
||||
|
||||
assert.True(t, b.Update(l, 200)) // We saw packet 55, 100, and 200 and can still track 190,191,192,193,194,195,196,197,198,199
|
||||
assert.Equal(t, int64(188), b.lostCounter.Count())
|
||||
// Update(200): same shape: 16 - 1 = 15 evicted unset, plus 200 - 100 - 16 = 84 past window. Total 99.
|
||||
assert.True(t, b.Update(l, 200))
|
||||
assert.Equal(t, int64(39+44+99), b.lostCounter.Count())
|
||||
}
|
||||
|
||||
func TestBitsDupeCounter(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(10)
|
||||
b := NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
b.dupeCounter.Clear()
|
||||
b.outOfWindowCounter.Clear()
|
||||
@@ -114,120 +139,117 @@ func TestBitsDupeCounter(t *testing.T) {
|
||||
|
||||
func TestBitsOutOfWindowCounter(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(10)
|
||||
b := NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
b.dupeCounter.Clear()
|
||||
b.outOfWindowCounter.Clear()
|
||||
|
||||
// Jump to 20 (warmup branch + 4 past-window packets).
|
||||
assert.True(t, b.Update(l, 20))
|
||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||
|
||||
assert.True(t, b.Update(l, 21))
|
||||
assert.True(t, b.Update(l, 22))
|
||||
assert.True(t, b.Update(l, 23))
|
||||
assert.True(t, b.Update(l, 24))
|
||||
assert.True(t, b.Update(l, 25))
|
||||
assert.True(t, b.Update(l, 26))
|
||||
assert.True(t, b.Update(l, 27))
|
||||
assert.True(t, b.Update(l, 28))
|
||||
assert.True(t, b.Update(l, 29))
|
||||
// 9 single-step advances, each evicts a slot whose bit was cleared during
|
||||
// the jump above and whose value was never seen, so each contributes 1
|
||||
// to lostCounter.
|
||||
for n := uint64(21); n <= 29; n++ {
|
||||
assert.True(t, b.Update(l, n))
|
||||
}
|
||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||
|
||||
// 0 is below current-length (29-16=13) so it falls outside the window.
|
||||
assert.False(t, b.Update(l, 0))
|
||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||
|
||||
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
||||
// 4 from the Update(20) jump + 9 from 21..29.
|
||||
assert.Equal(t, int64(13), b.lostCounter.Count())
|
||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||
}
|
||||
|
||||
func TestBitsLostCounter(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(10)
|
||||
b := NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
b.dupeCounter.Clear()
|
||||
b.outOfWindowCounter.Clear()
|
||||
|
||||
assert.True(t, b.Update(l, 20))
|
||||
assert.True(t, b.Update(l, 21))
|
||||
assert.True(t, b.Update(l, 22))
|
||||
assert.True(t, b.Update(l, 23))
|
||||
assert.True(t, b.Update(l, 24))
|
||||
assert.True(t, b.Update(l, 25))
|
||||
assert.True(t, b.Update(l, 26))
|
||||
assert.True(t, b.Update(l, 27))
|
||||
assert.True(t, b.Update(l, 28))
|
||||
assert.True(t, b.Update(l, 29))
|
||||
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
||||
// Walk 20..29 like the original, just with a bigger window. Same
|
||||
// reasoning as TestBitsOutOfWindowCounter: 4 past-window from Update(20),
|
||||
// then 9 more from the unit advances.
|
||||
for n := uint64(20); n <= 29; n++ {
|
||||
assert.True(t, b.Update(l, n))
|
||||
}
|
||||
assert.Equal(t, int64(13), b.lostCounter.Count())
|
||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||
|
||||
b = NewBits(10)
|
||||
b = NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
b.dupeCounter.Clear()
|
||||
b.outOfWindowCounter.Clear()
|
||||
|
||||
assert.True(t, b.Update(l, 9))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
// 10 will set 0 index, 0 was already set, no lost packets
|
||||
assert.True(t, b.Update(l, 10))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
// 11 will set 1 index, 1 was missed, we should see 1 packet lost
|
||||
assert.True(t, b.Update(l, 11))
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
// Now let's fill in the window, should end up with 8 lost packets
|
||||
assert.True(t, b.Update(l, 12))
|
||||
assert.True(t, b.Update(l, 13))
|
||||
assert.True(t, b.Update(l, 14))
|
||||
// Update(15) clears the warmup window (no lost), sets slot 15.
|
||||
assert.True(t, b.Update(l, 15))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
|
||||
// Update(16): slot 0 was already set (NewBits seeded it), and 16 is not
|
||||
// strictly > length, so nothing is recorded as lost.
|
||||
assert.True(t, b.Update(l, 16))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
|
||||
// Update(17): we jumped straight from 0 to 15, so slot 1 was cleared
|
||||
// (and never re-set). 17 > 16 is past warmup, so packet 1 is recorded lost.
|
||||
assert.True(t, b.Update(l, 17))
|
||||
assert.True(t, b.Update(l, 18))
|
||||
assert.True(t, b.Update(l, 19))
|
||||
assert.Equal(t, int64(8), b.lostCounter.Count())
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
|
||||
// Jump ahead by a window size
|
||||
assert.True(t, b.Update(l, 29))
|
||||
assert.Equal(t, int64(8), b.lostCounter.Count())
|
||||
// Now lets walk ahead normally through the window, the missed packets should fill in
|
||||
assert.True(t, b.Update(l, 30))
|
||||
assert.True(t, b.Update(l, 31))
|
||||
assert.True(t, b.Update(l, 32))
|
||||
assert.True(t, b.Update(l, 33))
|
||||
assert.True(t, b.Update(l, 34))
|
||||
assert.True(t, b.Update(l, 35))
|
||||
assert.True(t, b.Update(l, 36))
|
||||
assert.True(t, b.Update(l, 37))
|
||||
assert.True(t, b.Update(l, 38))
|
||||
// 39 packets tracked, 22 seen, 17 lost
|
||||
assert.Equal(t, int64(17), b.lostCounter.Count())
|
||||
// Fill in 18..30 in single steps. Each i evicts slot i%16. Slots 2..14
|
||||
// were all cleared during Update(15), and we never re-set any of them,
|
||||
// so each i in 18..30 is a fresh lost packet — 13 more.
|
||||
for n := uint64(18); n <= 30; n++ {
|
||||
assert.True(t, b.Update(l, n))
|
||||
}
|
||||
assert.Equal(t, int64(14), b.lostCounter.Count())
|
||||
|
||||
// Jump ahead by 2 windows, should have recording 1 full window missing
|
||||
assert.True(t, b.Update(l, 58))
|
||||
assert.Equal(t, int64(27), b.lostCounter.Count())
|
||||
// Now lets walk ahead normally through the window, the missed packets should fill in from this window
|
||||
assert.True(t, b.Update(l, 59))
|
||||
assert.True(t, b.Update(l, 60))
|
||||
assert.True(t, b.Update(l, 61))
|
||||
assert.True(t, b.Update(l, 62))
|
||||
assert.True(t, b.Update(l, 63))
|
||||
assert.True(t, b.Update(l, 64))
|
||||
assert.True(t, b.Update(l, 65))
|
||||
assert.True(t, b.Update(l, 66))
|
||||
assert.True(t, b.Update(l, 67))
|
||||
// 68 packets tracked, 32 seen, 36 missed
|
||||
assert.Equal(t, int64(36), b.lostCounter.Count())
|
||||
// Jump ahead by exactly one window size.
|
||||
assert.True(t, b.Update(l, 46))
|
||||
// end = min(46, 30+16) = 46, count = 16, all slots cleared. Before the
|
||||
// jump every slot 0..15 had been set (Update(15), (16), (17), 18..30),
|
||||
// so wasSet=16 and 46 == current+length means no past-window slack:
|
||||
// lost contribution = 0.
|
||||
assert.Equal(t, int64(14), b.lostCounter.Count())
|
||||
|
||||
// Walk 47..55. The Update(46) jump cleared every slot, so only slot 14
|
||||
// (for packet 46) is set when we start. Each subsequent unit step lands
|
||||
// on a slot that was cleared and is past warmup, so it counts as lost.
|
||||
// 9 more = 23.
|
||||
for n := uint64(47); n <= 55; n++ {
|
||||
assert.True(t, b.Update(l, n))
|
||||
}
|
||||
assert.Equal(t, int64(23), b.lostCounter.Count())
|
||||
|
||||
// Jump ahead by two windows: clears the window plus past-window loss.
|
||||
assert.True(t, b.Update(l, 87))
|
||||
// current=55, length=16. end = min(87, 71) = 71. count=16, all slots
|
||||
// cleared. Slots set before the clear are slots 14,15,0..7 (10 total).
|
||||
// Lost from clear = 16 - 10 = 6. Past window: 87 - 55 - 16 = 16. +22.
|
||||
assert.Equal(t, int64(45), b.lostCounter.Count())
|
||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||
}
|
||||
|
||||
func TestBitsLostCounterIssue1(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(10)
|
||||
b := NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
b.dupeCounter.Clear()
|
||||
b.outOfWindowCounter.Clear()
|
||||
|
||||
// Receive 4, backfill 1, then 9, 2, 3, 5, 6, 7 (skip 8), 10, 11, 14.
|
||||
// Then jump to 25 — slot 25%16=9 is being evicted, but it had been set
|
||||
// (we received packet 9), so no spurious lost increment. The original
|
||||
// regression was about double-counting a missing packet when its slot
|
||||
// got cleared on a jump. With the jump path now using clearRange's
|
||||
// word-level wasSet count, the same semantics hold.
|
||||
assert.True(t, b.Update(l, 4))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 1))
|
||||
@@ -244,7 +266,7 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 7))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
// assert.True(t, b.Update(l, 8))
|
||||
// Skip packet 8.
|
||||
assert.True(t, b.Update(l, 10))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 11))
|
||||
@@ -252,9 +274,23 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
||||
|
||||
assert.True(t, b.Update(l, 14))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
// Issue seems to be here, we reset missing packet 8 to false here and don't increment the lost counter
|
||||
assert.True(t, b.Update(l, 19))
|
||||
|
||||
// Jump to 25. With length=16, slot 25%16=9 corresponds to packet 9
|
||||
// (which we DID receive), so its bit is set and no lost++ from that
|
||||
// eviction. The trace below shows the only loss is packet 8.
|
||||
assert.True(t, b.Update(l, 25))
|
||||
// current was 14, i=25. end=min(25,30)=25. count=11. startPos=15.
|
||||
// steady? current=14<16, so warmup branch: per-bit n=15..25, count those
|
||||
// with !get(n) AND n>16. n=17..25 are >16. Among slots 17%16=1..25%16=9
|
||||
// did we set slots 1..9 (packets 1..9)? Yes for all but slot 8 (packet 8
|
||||
// was skipped). n=24 maps to slot 8 which is FALSE → lost++. All other
|
||||
// n in 17..25 map to slots that are set. n=16 is not strictly > 16. So
|
||||
// lost = 1.
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
|
||||
// Fill in 12, 13, 15, 16. Each is below current=25 (in-window). 16 must
|
||||
// recheck slot 0 — it was set by NewBits and then cleared by the
|
||||
// Update(25) jump, so 16 backfills cleanly.
|
||||
assert.True(t, b.Update(l, 12))
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 13))
|
||||
@@ -263,29 +299,140 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 16))
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 17))
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 18))
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 20))
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.True(t, b.Update(l, 21))
|
||||
|
||||
// We missed packet 8 above
|
||||
// We missed packet 8 above and that loss is still recorded once, never
|
||||
// double-counted, never zeroed.
|
||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||
}
|
||||
|
||||
func BenchmarkBits(b *testing.B) {
|
||||
z := NewBits(10)
|
||||
for n := 0; n < b.N; n++ {
|
||||
for i := range z.bits {
|
||||
z.bits[i] = true
|
||||
}
|
||||
for i := range z.bits {
|
||||
z.bits[i] = false
|
||||
}
|
||||
// TestBitsWarmupOvershoot exercises the jump path's warmup arm with an
|
||||
// overshoot past one full window. NewBits leaves current=0 with only slot 0
|
||||
// "set" by the marker. Jumping straight to length+k must (a) clear every
|
||||
// slot the jump straddles, (b) count only past-window slack (not the
|
||||
// in-window slots, which never had a "lost" tenant during warmup), and
|
||||
// (c) leave the cursor at the new counter so subsequent unit advances
|
||||
// count from steady state. The marker bit at slot 0 is irrelevant once
|
||||
// current >= length.
|
||||
func TestBitsWarmupOvershoot(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(16)
|
||||
b.lostCounter.Clear()
|
||||
|
||||
// Jump from current=0 to i=20 (length=16, overshoot=4).
|
||||
// Warmup arm: counts slots in [1..16] where bit unset and n>length.
|
||||
// Only n=16 was unset and >length: but slot 16%16=0 is the marker,
|
||||
// so b.get(16) reads bits[0]=1 and skips. Result: 0 lost from the loop.
|
||||
// Past-window: i - current - length = 20 - 0 - 16 = 4 lost.
|
||||
assert.True(t, b.Update(l, 20))
|
||||
assert.Equal(t, int64(4), b.lostCounter.Count())
|
||||
assert.Equal(t, uint64(20), b.current)
|
||||
|
||||
// Steady state now (current=20 >= length=16). Unit advance to 21
|
||||
// stomps slot 21%16=5, which was cleared by the jump and not reset,
|
||||
// so this is +1 lost.
|
||||
assert.True(t, b.Update(l, 21))
|
||||
assert.Equal(t, int64(5), b.lostCounter.Count())
|
||||
}
|
||||
|
||||
// TestBitsCheckAcrossWarmupBoundary pins the underflow trick in Check's
|
||||
// in-window clause. While in warmup, b.current-b.length underflows uint64
|
||||
// to a huge value so the first OR-clause is always false; the second
|
||||
// clause (i < length && current < length) carries the in-window check.
|
||||
// Once current >= length the regimes flip cleanly.
|
||||
func TestBitsCheckAcrossWarmupBoundary(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(16)
|
||||
|
||||
// Warmup: current=0. Check(0) must read the marker (set) and return false.
|
||||
assert.False(t, b.Check(l, 0), "marker slot should look already-received")
|
||||
// Warmup: any 0 < i < length is in-window and unset → accepted.
|
||||
for i := uint64(1); i < 16; i++ {
|
||||
assert.True(t, b.Check(l, i), "warmup in-window i=%d should be accepted", i)
|
||||
}
|
||||
// Warmup: i >= length but > current is "next number" so accepted.
|
||||
assert.True(t, b.Check(l, 16))
|
||||
assert.True(t, b.Check(l, 1_000_000))
|
||||
|
||||
// Cross into steady state.
|
||||
assert.True(t, b.Update(l, 100))
|
||||
// Now current=100, length=16. In-window range is [85..100].
|
||||
// 84 is just outside: the underflow clause activates; 84 > 100-16=84 is false.
|
||||
// And the warmup clause is false (current >= length). So out of window.
|
||||
assert.False(t, b.Check(l, 84))
|
||||
// 85 sits at the boundary. 85 > 84 is true → in window, unset → accept.
|
||||
assert.True(t, b.Check(l, 85))
|
||||
// 100 is current itself; not strictly greater, in-window, but already set.
|
||||
assert.False(t, b.Check(l, 100))
|
||||
// Way out: clearly out of window.
|
||||
assert.False(t, b.Check(l, 50))
|
||||
}
|
||||
|
||||
// TestBitsMarkerInvariant verifies the seeded bits[0]=1 marker behaves
|
||||
// correctly across warmup and beyond. Update should never clear the marker
|
||||
// during warmup (clearRange skips position 0 when startPos=1), and once
|
||||
// current >= length the marker is no longer consulted by Check/Update on
|
||||
// the live path — but it must still report counter 0 as a duplicate while
|
||||
// we are in warmup.
|
||||
func TestBitsMarkerInvariant(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
b := NewBits(8)
|
||||
|
||||
// Counter 0 is the seeded marker; Check sees it as already received.
|
||||
assert.False(t, b.Check(l, 0))
|
||||
// Update(0) at current=0 hits the duplicate branch.
|
||||
b.dupeCounter.Clear()
|
||||
assert.False(t, b.Update(l, 0))
|
||||
assert.Equal(t, int64(1), b.dupeCounter.Count())
|
||||
|
||||
// Walk forward through warmup; the marker must remain set.
|
||||
for n := uint64(1); n <= 7; n++ {
|
||||
assert.True(t, b.Update(l, n))
|
||||
}
|
||||
// Position 0 (the marker) should still read as set because we never
|
||||
// cleared it; Update(0) still looks like a duplicate.
|
||||
assert.False(t, b.Check(l, 0))
|
||||
|
||||
// Cross into steady state with a unit advance to 8: pos=0, evicts the
|
||||
// marker bit. The lost-counter guard (i > b.length) is false (8 == 8),
|
||||
// so this advance does NOT charge a lost packet — exactly what the
|
||||
// marker is there to prevent.
|
||||
b.lostCounter.Clear()
|
||||
assert.True(t, b.Update(l, 8))
|
||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||
// The slot at pos 0 is now occupied by counter 8.
|
||||
assert.False(t, b.Check(l, 8))
|
||||
}
|
||||
|
||||
// BenchmarkBitsUpdateInOrder is the steady-state hot path: each call is
|
||||
// i == current+1.
|
||||
func BenchmarkBitsUpdateInOrder(b *testing.B) {
|
||||
l := test.NewLogger()
|
||||
z := NewBits(16384)
|
||||
for n := 0; n < b.N; n++ {
|
||||
z.Update(l, uint64(n)+1)
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkBitsUpdateReorder simulates light reorder within the window:
|
||||
// every other packet arrives one slot behind its predecessor (forces the
|
||||
// in-window backfill branch).
|
||||
func BenchmarkBitsUpdateReorder(b *testing.B) {
|
||||
l := test.NewLogger()
|
||||
z := NewBits(16384)
|
||||
for n := 0; n < b.N; n++ {
|
||||
base := uint64(n) * 2
|
||||
z.Update(l, base+2)
|
||||
z.Update(l, base+1)
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkBitsUpdateLargeJumps stresses the clearRange word-level path.
|
||||
func BenchmarkBitsUpdateLargeJumps(b *testing.B) {
|
||||
l := test.NewLogger()
|
||||
z := NewBits(16384)
|
||||
for n := 0; n < b.N; n++ {
|
||||
z.Update(l, uint64(n+1)*1000)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -217,6 +217,10 @@ func (ncp *CAPool) verify(c Certificate, now time.Time, certFp string, signerFp
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if signer.Certificate.Curve() != c.Curve() {
|
||||
return nil, ErrCurveMismatch
|
||||
}
|
||||
|
||||
if signer.Certificate.Expired(now) {
|
||||
return nil, ErrRootExpired
|
||||
}
|
||||
|
||||
@@ -654,3 +654,31 @@ func TestCertificateV2_Verify_Subnets(t *testing.T) {
|
||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestCertificateV2_CurveMismatch(t *testing.T) {
|
||||
caIp1 := mustParsePrefixUnmapped("10.0.0.0/16")
|
||||
caIp2 := mustParsePrefixUnmapped("192.168.0.0/24")
|
||||
ca, _, caKey, _ := NewTestCaCert(Version2, Curve_P256, time.Now(), time.Now().Add(10*time.Minute), []netip.Prefix{caIp1, caIp2}, nil, []string{"test"})
|
||||
|
||||
caPem, err := ca.MarshalPEM()
|
||||
require.NoError(t, err)
|
||||
|
||||
caPool := NewCAPool()
|
||||
b, err := caPool.AddCAFromPEM(caPem)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, b)
|
||||
|
||||
// ip is outside the network
|
||||
cIp1 := mustParsePrefixUnmapped("10.0.0.1/24")
|
||||
c, _, _, _ := NewTestCert(Version2, Curve_P256, ca, caKey, "test", time.Now(), time.Now().Add(5*time.Minute), []netip.Prefix{cIp1}, nil, []string{"test"})
|
||||
|
||||
fp, _ := c.Fingerprint()
|
||||
_, err = caPool.verify(c, time.Now(), fp, c.Issuer())
|
||||
require.NoError(t, err)
|
||||
//
|
||||
c2 := c.(*certificateV2)
|
||||
c2.curve = Curve_CURVE25519
|
||||
fp, _ = c.Fingerprint()
|
||||
_, err = caPool.verify(c, time.Now(), fp, c.Issuer())
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
@@ -112,6 +112,9 @@ func (c *certificateV1) CheckSignature(key []byte) bool {
|
||||
}
|
||||
switch c.details.curve {
|
||||
case Curve_CURVE25519:
|
||||
if len(key) != ed25519.PublicKeySize {
|
||||
return false //avoids a panic internal to ed25519
|
||||
}
|
||||
return ed25519.Verify(key, b, c.signature)
|
||||
case Curve_P256:
|
||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
||||
|
||||
@@ -151,6 +151,9 @@ func (c *certificateV2) CheckSignature(key []byte) bool {
|
||||
|
||||
switch c.curve {
|
||||
case Curve_CURVE25519:
|
||||
if len(key) != ed25519.PublicKeySize {
|
||||
return false //avoids a panic internal to ed25519
|
||||
}
|
||||
return ed25519.Verify(key, b, c.signature)
|
||||
case Curve_P256:
|
||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
||||
|
||||
@@ -22,6 +22,7 @@ var (
|
||||
ErrCaNotFound = errors.New("could not find ca for the certificate")
|
||||
ErrUnknownVersion = errors.New("certificate version unrecognized")
|
||||
ErrCertPubkeyPresent = errors.New("certificate has unexpected pubkey present")
|
||||
ErrCurveMismatch = errors.New("certificate curve does not match CA")
|
||||
|
||||
ErrInvalidPEMBlock = errors.New("input did not contain a valid PEM encoded block")
|
||||
ErrInvalidPEMCertificateBanner = errors.New("bytes did not contain a proper certificate banner")
|
||||
|
||||
@@ -163,3 +163,55 @@ func P256Keypair() ([]byte, []byte) {
|
||||
pubkey := privkey.PublicKey()
|
||||
return pubkey.Bytes(), privkey.Bytes()
|
||||
}
|
||||
|
||||
// DummyCert is a minimal cert.Certificate implementation for testing error paths.
|
||||
type DummyCert struct {
|
||||
Version_ cert.Version
|
||||
Curve_ cert.Curve
|
||||
Groups_ []string
|
||||
IsCA_ bool
|
||||
Issuer_ string
|
||||
Name_ string
|
||||
Networks_ []netip.Prefix
|
||||
NotAfter_ time.Time
|
||||
NotBefore_ time.Time
|
||||
PublicKey_ []byte
|
||||
Signature_ []byte
|
||||
UnsafeNetworks_ []netip.Prefix
|
||||
}
|
||||
|
||||
func (d *DummyCert) Version() cert.Version { return d.Version_ }
|
||||
func (d *DummyCert) Curve() cert.Curve { return d.Curve_ }
|
||||
func (d *DummyCert) Groups() []string { return d.Groups_ }
|
||||
func (d *DummyCert) IsCA() bool { return d.IsCA_ }
|
||||
func (d *DummyCert) Issuer() string { return d.Issuer_ }
|
||||
func (d *DummyCert) Name() string { return d.Name_ }
|
||||
func (d *DummyCert) Networks() []netip.Prefix { return d.Networks_ }
|
||||
func (d *DummyCert) NotAfter() time.Time { return d.NotAfter_ }
|
||||
func (d *DummyCert) NotBefore() time.Time { return d.NotBefore_ }
|
||||
func (d *DummyCert) PublicKey() []byte { return d.PublicKey_ }
|
||||
func (d *DummyCert) Signature() []byte { return d.Signature_ }
|
||||
func (d *DummyCert) UnsafeNetworks() []netip.Prefix { return d.UnsafeNetworks_ }
|
||||
func (d *DummyCert) Fingerprint() (string, error) { return "", nil }
|
||||
func (d *DummyCert) CheckSignature(key []byte) bool { return false }
|
||||
func (d *DummyCert) MarshalForHandshakes() ([]byte, error) { return nil, nil }
|
||||
func (d *DummyCert) MarshalPEM() ([]byte, error) { return nil, nil }
|
||||
func (d *DummyCert) MarshalJSON() ([]byte, error) { return nil, nil }
|
||||
func (d *DummyCert) Marshal() ([]byte, error) { return nil, nil }
|
||||
func (d *DummyCert) String() string { return "dummy" }
|
||||
func (d *DummyCert) Copy() cert.Certificate { return d }
|
||||
func (d *DummyCert) VerifyPrivateKey(c cert.Curve, k []byte) error { return nil }
|
||||
func (d *DummyCert) Expired(time.Time) bool { return false }
|
||||
func (d *DummyCert) MarshalPublicKeyPEM() []byte { return nil }
|
||||
func (d *DummyCert) PublicKeyPEM() []byte { return nil }
|
||||
|
||||
// NewTestCAPool creates a CAPool from the given CA certificates, panicking on error.
|
||||
func NewTestCAPool(cas ...cert.Certificate) *cert.CAPool {
|
||||
pool := cert.NewCAPool()
|
||||
for _, ca := range cas {
|
||||
if err := pool.AddCA(ca); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
return pool
|
||||
}
|
||||
|
||||
@@ -3,8 +3,15 @@
|
||||
|
||||
package main
|
||||
|
||||
import "github.com/sirupsen/logrus"
|
||||
import (
|
||||
"log/slog"
|
||||
"os"
|
||||
|
||||
func HookLogger(l *logrus.Logger) {
|
||||
// Do nothing, let the logs flow to stdout/stderr
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
// newPlatformLogger returns a *slog.Logger that writes to stdout. Non-Windows
|
||||
// platforms have no special sink to integrate with.
|
||||
func newPlatformLogger() *slog.Logger {
|
||||
return logging.NewLogger(os.Stdout)
|
||||
}
|
||||
|
||||
@@ -1,54 +1,86 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io/ioutil"
|
||||
"os"
|
||||
"context"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
// HookLogger routes the logrus logs through the service logger so that they end up in the Windows Event Viewer
|
||||
// logrus output will be discarded
|
||||
func HookLogger(l *logrus.Logger) {
|
||||
l.AddHook(newLogHook(logger))
|
||||
l.SetOutput(ioutil.Discard)
|
||||
// newPlatformLogger returns a *slog.Logger that routes every log record
|
||||
// through the Windows service logger so records end up in the Windows
|
||||
// Event Log. All the heavy lifting (level management, format swap,
|
||||
// timestamp toggle, WithAttrs/WithGroup) comes from logging.NewHandler;
|
||||
// this file only contributes:
|
||||
//
|
||||
// - an io.Writer that forwards each formatted line to the service
|
||||
// logger at the current record's Event Log severity, and
|
||||
// - a thin severityTag that embeds *logging.Handler and overrides
|
||||
// only Handle / WithAttrs / WithGroup, so Event Viewer's severity
|
||||
// column and severity-based filters keep working the way they did
|
||||
// before the slog migration.
|
||||
//
|
||||
// Format (text vs json) is carried by the embedded *logging.Handler, so
|
||||
// logging.format: json in config still produces JSON lines in Event
|
||||
// Viewer, same as the pre-slog logrus setup.
|
||||
func newPlatformLogger() *slog.Logger {
|
||||
w := &eventLogWriter{}
|
||||
return slog.New(&severityTag{Handler: logging.NewHandler(w), w: w})
|
||||
}
|
||||
|
||||
type logHook struct {
|
||||
sl service.Logger
|
||||
// eventLogWriter forwards slog-formatted lines to the Windows service
|
||||
// logger at the severity most recently stashed by severityTag.Handle.
|
||||
// The mutex serializes the stash + inner.Handle + Write cycle per record
|
||||
// across all concurrent goroutines; slog's builtin text/json handlers
|
||||
// each hold their own mutex around Write, but that only protects the
|
||||
// Write call itself, not our stash-then-handle sequence.
|
||||
type eventLogWriter struct {
|
||||
mu sync.Mutex
|
||||
level slog.Level
|
||||
}
|
||||
|
||||
func newLogHook(sl service.Logger) *logHook {
|
||||
return &logHook{sl: sl}
|
||||
}
|
||||
|
||||
func (h *logHook) Fire(entry *logrus.Entry) error {
|
||||
line, err := entry.String()
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "Unable to read entry, %v", err)
|
||||
return err
|
||||
}
|
||||
|
||||
switch entry.Level {
|
||||
case logrus.PanicLevel:
|
||||
return h.sl.Error(line)
|
||||
case logrus.FatalLevel:
|
||||
return h.sl.Error(line)
|
||||
case logrus.ErrorLevel:
|
||||
return h.sl.Error(line)
|
||||
case logrus.WarnLevel:
|
||||
return h.sl.Warning(line)
|
||||
case logrus.InfoLevel:
|
||||
return h.sl.Info(line)
|
||||
case logrus.DebugLevel:
|
||||
return h.sl.Info(line)
|
||||
func (w *eventLogWriter) Write(p []byte) (int, error) {
|
||||
line := strings.TrimRight(string(p), "\n")
|
||||
switch {
|
||||
case w.level >= slog.LevelError:
|
||||
return len(p), logger.Error(line)
|
||||
case w.level >= slog.LevelWarn:
|
||||
return len(p), logger.Warning(line)
|
||||
default:
|
||||
return nil
|
||||
return len(p), logger.Info(line)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *logHook) Levels() []logrus.Level {
|
||||
return logrus.AllLevels
|
||||
// severityTag embeds *logging.Handler to pick up everything it does for
|
||||
// free (Enabled, SetLevel, GetLevel, SetFormat, GetFormat,
|
||||
// SetDisableTimestamp) and overrides only Handle / WithAttrs / WithGroup
|
||||
// so each record's slog.Level is stashed on the writer before formatting
|
||||
// and so derived handlers stay wrapped as severityTag rather than
|
||||
// downgrading to bare *logging.Handler.
|
||||
type severityTag struct {
|
||||
*logging.Handler
|
||||
w *eventLogWriter
|
||||
}
|
||||
|
||||
func (s *severityTag) Handle(ctx context.Context, r slog.Record) error {
|
||||
s.w.mu.Lock()
|
||||
defer s.w.mu.Unlock()
|
||||
s.w.level = r.Level
|
||||
return s.Handler.Handle(ctx, r)
|
||||
}
|
||||
|
||||
func (s *severityTag) WithAttrs(attrs []slog.Attr) slog.Handler {
|
||||
if len(attrs) == 0 {
|
||||
return s
|
||||
}
|
||||
return &severityTag{Handler: s.Handler.WithAttrs(attrs).(*logging.Handler), w: s.w}
|
||||
}
|
||||
|
||||
func (s *severityTag) WithGroup(name string) slog.Handler {
|
||||
if name == "" {
|
||||
return s
|
||||
}
|
||||
return &severityTag{Handler: s.Handler.WithGroup(name).(*logging.Handler), w: s.w}
|
||||
}
|
||||
|
||||
@@ -7,9 +7,9 @@ import (
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
|
||||
@@ -50,9 +50,14 @@ func main() {
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
l := logging.NewLogger(os.Stdout)
|
||||
|
||||
if *serviceFlag != "" {
|
||||
doService(configPath, configTest, Build, serviceFlag)
|
||||
os.Exit(1)
|
||||
if err := doService(configPath, configTest, Build, serviceFlag); err != nil {
|
||||
l.Error("Service command failed", "error", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if *configPath == "" {
|
||||
@@ -61,9 +66,6 @@ func main() {
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
l := logrus.New()
|
||||
l.Out = os.Stdout
|
||||
|
||||
c := config.NewC(l)
|
||||
err := c.Load(*configPath)
|
||||
if err != nil {
|
||||
@@ -71,6 +73,16 @@ func main() {
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
if err := logging.ApplyConfig(l, c); err != nil {
|
||||
fmt.Printf("failed to apply logging config: %s", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
if err := logging.ApplyConfig(l, c); err != nil {
|
||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
||||
}
|
||||
})
|
||||
|
||||
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
||||
if err != nil {
|
||||
util.LogWithContextIfNeeded("Failed to start", err, l)
|
||||
@@ -78,8 +90,20 @@ func main() {
|
||||
}
|
||||
|
||||
if !*configTest {
|
||||
ctrl.Start()
|
||||
ctrl.ShutdownBlock()
|
||||
wait, err := ctrl.Start()
|
||||
if err != nil {
|
||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
go ctrl.ShutdownBlock()
|
||||
|
||||
if err := wait(); err != nil {
|
||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||
os.Exit(2)
|
||||
}
|
||||
|
||||
l.Info("Goodbye")
|
||||
}
|
||||
|
||||
os.Exit(0)
|
||||
|
||||
@@ -7,9 +7,9 @@ import (
|
||||
"path/filepath"
|
||||
|
||||
"github.com/kardianos/service"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
var logger service.Logger
|
||||
@@ -25,8 +25,7 @@ func (p *program) Start(s service.Service) error {
|
||||
// Start should not block.
|
||||
logger.Info("Nebula service starting.")
|
||||
|
||||
l := logrus.New()
|
||||
HookLogger(l)
|
||||
l := newPlatformLogger()
|
||||
|
||||
c := config.NewC(l)
|
||||
err := c.Load(*p.configPath)
|
||||
@@ -34,6 +33,15 @@ func (p *program) Start(s service.Service) error {
|
||||
return fmt.Errorf("failed to load config: %s", err)
|
||||
}
|
||||
|
||||
if err := logging.ApplyConfig(l, c); err != nil {
|
||||
return fmt.Errorf("failed to apply logging config: %s", err)
|
||||
}
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
if err := logging.ApplyConfig(l, c); err != nil {
|
||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
||||
}
|
||||
})
|
||||
|
||||
p.control, err = nebula.Main(c, *p.configTest, Build, l, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -57,11 +65,11 @@ func fileExists(filename string) bool {
|
||||
return true
|
||||
}
|
||||
|
||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) {
|
||||
func doService(configPath *string, configTest *bool, build string, serviceFlag *string) error {
|
||||
if *configPath == "" {
|
||||
ex, err := os.Executable()
|
||||
if err != nil {
|
||||
panic(err)
|
||||
return err
|
||||
}
|
||||
*configPath = filepath.Dir(ex) + "/config.yaml"
|
||||
if !fileExists(*configPath) {
|
||||
@@ -85,16 +93,16 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
||||
// Here are what the different loggers are doing:
|
||||
// - `log` is the standard go log utility, meant to be used while the process is still attached to stdout/stderr
|
||||
// - `logger` is the service log utility that may be attached to a special place depending on OS (Windows will have it attached to the event log)
|
||||
// - above, in `Run` we create a `logrus.Logger` which is what nebula expects to use
|
||||
// - in program.Start we build a *slog.Logger via newPlatformLogger; on non-Windows that is a stdout-backed slog logger, on Windows it routes records through the service logger
|
||||
s, err := service.New(prg, svcConfig)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
return err
|
||||
}
|
||||
|
||||
errs := make(chan error, 5)
|
||||
logger, err = s.Logger(errs)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
return err
|
||||
}
|
||||
|
||||
go func() {
|
||||
@@ -109,18 +117,16 @@ func doService(configPath *string, configTest *bool, build string, serviceFlag *
|
||||
|
||||
switch *serviceFlag {
|
||||
case "run":
|
||||
err = s.Run()
|
||||
if err != nil {
|
||||
if err := s.Run(); err != nil {
|
||||
// Route any errors to the system logger
|
||||
logger.Error(err)
|
||||
}
|
||||
default:
|
||||
err := service.Control(s, *serviceFlag)
|
||||
if err != nil {
|
||||
if err := service.Control(s, *serviceFlag); err != nil {
|
||||
log.Printf("Valid actions: %q\n", service.ControlAction)
|
||||
log.Fatal(err)
|
||||
return err
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
+26
-5
@@ -7,9 +7,9 @@ import (
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
|
||||
@@ -55,8 +55,7 @@ func main() {
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
l := logrus.New()
|
||||
l.Out = os.Stdout
|
||||
l := logging.NewLogger(os.Stdout)
|
||||
|
||||
c := config.NewC(l)
|
||||
err := c.Load(*configPath)
|
||||
@@ -65,6 +64,16 @@ func main() {
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
if err := logging.ApplyConfig(l, c); err != nil {
|
||||
fmt.Printf("failed to apply logging config: %s", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
if err := logging.ApplyConfig(l, c); err != nil {
|
||||
l.Error("Failed to reconfigure logger on reload", "error", err)
|
||||
}
|
||||
})
|
||||
|
||||
ctrl, err := nebula.Main(c, *configTest, Build, l, nil)
|
||||
if err != nil {
|
||||
util.LogWithContextIfNeeded("Failed to start", err, l)
|
||||
@@ -72,9 +81,21 @@ func main() {
|
||||
}
|
||||
|
||||
if !*configTest {
|
||||
ctrl.Start()
|
||||
wait, err := ctrl.Start()
|
||||
if err != nil {
|
||||
util.LogWithContextIfNeeded("Error while running", err, l)
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
go ctrl.ShutdownBlock()
|
||||
notifyReady(l)
|
||||
ctrl.ShutdownBlock()
|
||||
|
||||
if err := wait(); err != nil {
|
||||
l.Error("Nebula stopped due to fatal error", "error", err)
|
||||
os.Exit(2)
|
||||
}
|
||||
|
||||
l.Info("Goodbye")
|
||||
}
|
||||
|
||||
os.Exit(0)
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// SdNotifyReady tells systemd the service is ready and dependent services can now be started
|
||||
@@ -13,30 +12,30 @@ import (
|
||||
// https://www.freedesktop.org/software/systemd/man/systemd.service.html
|
||||
const SdNotifyReady = "READY=1"
|
||||
|
||||
func notifyReady(l *logrus.Logger) {
|
||||
func notifyReady(l *slog.Logger) {
|
||||
sockName := os.Getenv("NOTIFY_SOCKET")
|
||||
if sockName == "" {
|
||||
l.Debugln("NOTIFY_SOCKET systemd env var not set, not sending ready signal")
|
||||
l.Debug("NOTIFY_SOCKET systemd env var not set, not sending ready signal")
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := net.DialTimeout("unixgram", sockName, time.Second)
|
||||
if err != nil {
|
||||
l.WithError(err).Error("failed to connect to systemd notification socket")
|
||||
l.Error("failed to connect to systemd notification socket", "error", err)
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
err = conn.SetWriteDeadline(time.Now().Add(time.Second))
|
||||
if err != nil {
|
||||
l.WithError(err).Error("failed to set the write deadline for the systemd notification socket")
|
||||
l.Error("failed to set the write deadline for the systemd notification socket", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
if _, err = conn.Write([]byte(SdNotifyReady)); err != nil {
|
||||
l.WithError(err).Error("failed to signal the systemd notification socket")
|
||||
l.Error("failed to signal the systemd notification socket", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
l.Debugln("notified systemd the service is ready")
|
||||
l.Debug("notified systemd the service is ready")
|
||||
}
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
|
||||
package main
|
||||
|
||||
import "github.com/sirupsen/logrus"
|
||||
import "log/slog"
|
||||
|
||||
func notifyReady(_ *logrus.Logger) {
|
||||
func notifyReady(_ *slog.Logger) {
|
||||
// No init service to notify
|
||||
}
|
||||
|
||||
+15
-6
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math"
|
||||
"os"
|
||||
"os/signal"
|
||||
@@ -16,7 +17,6 @@ import (
|
||||
"time"
|
||||
|
||||
"dario.cat/mergo"
|
||||
"github.com/sirupsen/logrus"
|
||||
"go.yaml.in/yaml/v3"
|
||||
)
|
||||
|
||||
@@ -26,11 +26,11 @@ type C struct {
|
||||
Settings map[string]any
|
||||
oldSettings map[string]any
|
||||
callbacks []func(*C)
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
reloadLock sync.Mutex
|
||||
}
|
||||
|
||||
func NewC(l *logrus.Logger) *C {
|
||||
func NewC(l *slog.Logger) *C {
|
||||
return &C{
|
||||
Settings: make(map[string]any),
|
||||
l: l,
|
||||
@@ -107,12 +107,18 @@ func (c *C) HasChanged(k string) bool {
|
||||
|
||||
newVals, err := yaml.Marshal(nv)
|
||||
if err != nil {
|
||||
c.l.WithField("config_path", k).WithError(err).Error("Error while marshaling new config")
|
||||
c.l.Error("Error while marshaling new config",
|
||||
"config_path", k,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
|
||||
oldVals, err := yaml.Marshal(ov)
|
||||
if err != nil {
|
||||
c.l.WithField("config_path", k).WithError(err).Error("Error while marshaling old config")
|
||||
c.l.Error("Error while marshaling old config",
|
||||
"config_path", k,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
|
||||
return string(newVals) != string(oldVals)
|
||||
@@ -154,7 +160,10 @@ func (c *C) ReloadConfig() {
|
||||
|
||||
err := c.Load(c.path)
|
||||
if err != nil {
|
||||
c.l.WithField("config_path", c.path).WithError(err).Error("Error occurred while reloading config")
|
||||
c.l.Error("Error occurred while reloading config",
|
||||
"config_path", c.path,
|
||||
"error", err,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
+77
-99
@@ -5,13 +5,12 @@ import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
@@ -45,19 +44,16 @@ type connectionManager struct {
|
||||
inactivityTimeout atomic.Int64
|
||||
dropInactive atomic.Bool
|
||||
|
||||
metricsTxPunchy metrics.Counter
|
||||
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func newConnectionManagerFromConfig(l *logrus.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
||||
func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
||||
cm := &connectionManager{
|
||||
hostMap: hm,
|
||||
l: l,
|
||||
punchy: p,
|
||||
relayUsed: make(map[uint32]struct{}),
|
||||
relayUsedLock: &sync.RWMutex{},
|
||||
metricsTxPunchy: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
||||
hostMap: hm,
|
||||
l: l,
|
||||
punchy: p,
|
||||
relayUsed: make(map[uint32]struct{}),
|
||||
relayUsedLock: &sync.RWMutex{},
|
||||
}
|
||||
|
||||
cm.reload(c, true)
|
||||
@@ -85,9 +81,10 @@ func (cm *connectionManager) reload(c *config.C, initial bool) {
|
||||
old := cm.getInactivityTimeout()
|
||||
cm.inactivityTimeout.Store((int64)(c.GetDuration("tunnels.inactivity_timeout", 10*time.Minute)))
|
||||
if !initial {
|
||||
cm.l.WithField("oldDuration", old).
|
||||
WithField("newDuration", cm.getInactivityTimeout()).
|
||||
Info("Inactivity timeout has changed")
|
||||
cm.l.Info("Inactivity timeout has changed",
|
||||
"oldDuration", old,
|
||||
"newDuration", cm.getInactivityTimeout(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -95,9 +92,10 @@ func (cm *connectionManager) reload(c *config.C, initial bool) {
|
||||
old := cm.dropInactive.Load()
|
||||
cm.dropInactive.Store(c.GetBool("tunnels.drop_inactive", false))
|
||||
if !initial {
|
||||
cm.l.WithField("oldBool", old).
|
||||
WithField("newBool", cm.dropInactive.Load()).
|
||||
Info("Drop inactive setting has changed")
|
||||
cm.l.Info("Drop inactive setting has changed",
|
||||
"oldBool", old,
|
||||
"newBool", cm.dropInactive.Load(),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -256,7 +254,7 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
||||
var err error
|
||||
index, err = AddRelay(cm.l, newhostinfo, cm.hostMap, r.PeerAddr, nil, r.Type, Requested)
|
||||
if err != nil {
|
||||
cm.l.WithError(err).Error("failed to migrate relay to new hostinfo")
|
||||
cm.l.Error("failed to migrate relay to new hostinfo", "error", err)
|
||||
continue
|
||||
}
|
||||
switch r.Type {
|
||||
@@ -304,16 +302,16 @@ func (cm *connectionManager) migrateRelayUsed(oldhostinfo, newhostinfo *HostInfo
|
||||
|
||||
msg, err := req.Marshal()
|
||||
if err != nil {
|
||||
cm.l.WithError(err).Error("failed to marshal Control message to migrate relay")
|
||||
cm.l.Error("failed to marshal Control message to migrate relay", "error", err)
|
||||
} else {
|
||||
cm.intf.SendMessageToHostInfo(header.Control, 0, newhostinfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||
cm.l.WithFields(logrus.Fields{
|
||||
"relayFrom": req.RelayFromAddr,
|
||||
"relayTo": req.RelayToAddr,
|
||||
"initiatorRelayIndex": req.InitiatorRelayIndex,
|
||||
"responderRelayIndex": req.ResponderRelayIndex,
|
||||
"vpnAddrs": newhostinfo.vpnAddrs}).
|
||||
Info("send CreateRelayRequest")
|
||||
cm.l.Info("send CreateRelayRequest",
|
||||
"relayFrom", req.RelayFromAddr,
|
||||
"relayTo", req.RelayToAddr,
|
||||
"initiatorRelayIndex", req.InitiatorRelayIndex,
|
||||
"responderRelayIndex", req.ResponderRelayIndex,
|
||||
"vpnAddrs", newhostinfo.vpnAddrs,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -325,7 +323,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
|
||||
hostinfo := cm.hostMap.Indexes[localIndex]
|
||||
if hostinfo == nil {
|
||||
cm.l.WithField("localIndex", localIndex).Debugln("Not found in hostmap")
|
||||
cm.l.Debug("Not found in hostmap", "localIndex", localIndex)
|
||||
return doNothing, nil, nil
|
||||
}
|
||||
|
||||
@@ -345,10 +343,10 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
// A hostinfo is determined alive if there is incoming traffic
|
||||
if inTraffic {
|
||||
decision := doNothing
|
||||
if cm.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(cm.l).
|
||||
WithField("tunnelCheck", m{"state": "alive", "method": "passive"}).
|
||||
Debug("Tunnel status")
|
||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
||||
"tunnelCheck", m{"state": "alive", "method": "passive"},
|
||||
)
|
||||
}
|
||||
hostinfo.pendingDeletion.Store(false)
|
||||
|
||||
@@ -367,7 +365,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
|
||||
if !outTraffic {
|
||||
// Send a punch packet to keep the NAT state alive
|
||||
cm.sendPunch(hostinfo)
|
||||
cm.punchy.SendPunch(hostinfo)
|
||||
}
|
||||
|
||||
return decision, hostinfo, primary
|
||||
@@ -375,9 +373,9 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
|
||||
if hostinfo.pendingDeletion.Load() {
|
||||
// We have already sent a test packet and nothing was returned, this hostinfo is dead
|
||||
hostinfo.logger(cm.l).
|
||||
WithField("tunnelCheck", m{"state": "dead", "method": "active"}).
|
||||
Info("Tunnel status")
|
||||
hostinfo.logger(cm.l).Info("Tunnel status",
|
||||
"tunnelCheck", m{"state": "dead", "method": "active"},
|
||||
)
|
||||
|
||||
return deleteTunnel, hostinfo, nil
|
||||
}
|
||||
@@ -388,40 +386,39 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
||||
inactiveFor, isInactive := cm.isInactive(hostinfo, now)
|
||||
if isInactive {
|
||||
// Tunnel is inactive, tear it down
|
||||
hostinfo.logger(cm.l).
|
||||
WithField("inactiveDuration", inactiveFor).
|
||||
WithField("primary", mainHostInfo).
|
||||
Info("Dropping tunnel due to inactivity")
|
||||
hostinfo.logger(cm.l).Info("Dropping tunnel due to inactivity",
|
||||
"inactiveDuration", inactiveFor,
|
||||
"primary", mainHostInfo,
|
||||
)
|
||||
|
||||
return closeTunnel, hostinfo, primary
|
||||
}
|
||||
|
||||
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
||||
// Just maintain NAT state if configured to do so.
|
||||
cm.sendPunch(hostinfo)
|
||||
cm.punchy.SendPunch(hostinfo)
|
||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
||||
return doNothing, nil, nil
|
||||
}
|
||||
|
||||
if cm.punchy.GetTargetEverything() {
|
||||
// This is similar to the old punchy behavior with a slight optimization.
|
||||
// We aren't receiving traffic but we are sending it, punch on all known
|
||||
// ips in case we need to re-prime NAT state
|
||||
cm.sendPunch(hostinfo)
|
||||
}
|
||||
// We aren't receiving traffic but we are sending it. The outbound
|
||||
// traffic itself refreshes the primary remote's NAT state; this
|
||||
// fans out to non-primary remotes, but only if target_all_remotes
|
||||
// is configured.
|
||||
cm.punchy.SendPunchToAll(hostinfo)
|
||||
|
||||
if cm.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(cm.l).
|
||||
WithField("tunnelCheck", m{"state": "testing", "method": "active"}).
|
||||
Debug("Tunnel status")
|
||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
||||
"tunnelCheck", m{"state": "testing", "method": "active"},
|
||||
)
|
||||
}
|
||||
|
||||
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
||||
decision = sendTestPacket
|
||||
|
||||
} else {
|
||||
if cm.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(cm.l).Debugf("Hostinfo sadness")
|
||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(cm.l).Debug("Hostinfo sadness")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -493,14 +490,16 @@ func (cm *connectionManager) isInvalidCertificate(now time.Time, hostinfo *HostI
|
||||
return false //cert is still valid! yay!
|
||||
} else if err == cert.ErrBlockListed { //avoiding errors.Is for speed
|
||||
// Block listed certificates should always be disconnected
|
||||
hostinfo.logger(cm.l).WithError(err).
|
||||
WithField("fingerprint", remoteCert.Fingerprint).
|
||||
Info("Remote certificate is blocked, tearing down the tunnel")
|
||||
hostinfo.logger(cm.l).Info("Remote certificate is blocked, tearing down the tunnel",
|
||||
"error", err,
|
||||
"fingerprint", remoteCert.Fingerprint,
|
||||
)
|
||||
return true
|
||||
} else if cm.intf.disconnectInvalid.Load() {
|
||||
hostinfo.logger(cm.l).WithError(err).
|
||||
WithField("fingerprint", remoteCert.Fingerprint).
|
||||
Info("Remote certificate is no longer valid, tearing down the tunnel")
|
||||
hostinfo.logger(cm.l).Info("Remote certificate is no longer valid, tearing down the tunnel",
|
||||
"error", err,
|
||||
"fingerprint", remoteCert.Fingerprint,
|
||||
)
|
||||
return true
|
||||
} else {
|
||||
//if we reach here, the cert is no longer valid, but we're configured to keep tunnels from now-invalid certs open
|
||||
@@ -508,41 +507,17 @@ func (cm *connectionManager) isInvalidCertificate(now time.Time, hostinfo *HostI
|
||||
}
|
||||
}
|
||||
|
||||
func (cm *connectionManager) sendPunch(hostinfo *HostInfo) {
|
||||
if !cm.punchy.GetPunch() {
|
||||
// Punching is disabled
|
||||
return
|
||||
}
|
||||
|
||||
if cm.intf.lightHouse.IsAnyLighthouseAddr(hostinfo.vpnAddrs) {
|
||||
// Do not punch to lighthouses, we assume our lighthouse update interval is good enough.
|
||||
// In the event the update interval is not sufficient to maintain NAT state then a publicly available lighthouse
|
||||
// would lose the ability to notify us and punchy.respond would become unreliable.
|
||||
return
|
||||
}
|
||||
|
||||
if cm.punchy.GetTargetEverything() {
|
||||
hostinfo.remotes.ForEach(cm.hostMap.GetPreferredRanges(), func(addr netip.AddrPort, preferred bool) {
|
||||
cm.metricsTxPunchy.Inc(1)
|
||||
cm.intf.outside.WriteTo([]byte{1}, addr)
|
||||
})
|
||||
|
||||
} else if hostinfo.remote.IsValid() {
|
||||
cm.metricsTxPunchy.Inc(1)
|
||||
cm.intf.outside.WriteTo([]byte{1}, hostinfo.remote)
|
||||
}
|
||||
}
|
||||
|
||||
func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||
cs := cm.intf.pki.getCertState()
|
||||
curCrt := hostinfo.ConnectionState.myCert
|
||||
curCrtVersion := curCrt.Version()
|
||||
myCrt := cs.getCertificate(curCrtVersion)
|
||||
if myCrt == nil {
|
||||
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("version", curCrtVersion).
|
||||
WithField("reason", "local certificate removed").
|
||||
Info("Re-handshaking with remote")
|
||||
cm.l.Info("Re-handshaking with remote",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
"version", curCrtVersion,
|
||||
"reason", "local certificate removed",
|
||||
)
|
||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||
return
|
||||
}
|
||||
@@ -550,11 +525,12 @@ func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||
if peerCrt != nil && curCrtVersion < peerCrt.Certificate.Version() {
|
||||
// if our certificate version is less than theirs, and we have a matching version available, rehandshake?
|
||||
if cs.getCertificate(peerCrt.Certificate.Version()) != nil {
|
||||
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("version", curCrtVersion).
|
||||
WithField("peerVersion", peerCrt.Certificate.Version()).
|
||||
WithField("reason", "local certificate version lower than peer, attempting to correct").
|
||||
Info("Re-handshaking with remote")
|
||||
cm.l.Info("Re-handshaking with remote",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
"version", curCrtVersion,
|
||||
"peerVersion", peerCrt.Certificate.Version(),
|
||||
"reason", "local certificate version lower than peer, attempting to correct",
|
||||
)
|
||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(hh *HandshakeHostInfo) {
|
||||
hh.initiatingVersionOverride = peerCrt.Certificate.Version()
|
||||
})
|
||||
@@ -562,17 +538,19 @@ func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||
}
|
||||
}
|
||||
if !bytes.Equal(curCrt.Signature(), myCrt.Signature()) {
|
||||
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("reason", "local certificate is not current").
|
||||
Info("Re-handshaking with remote")
|
||||
cm.l.Info("Re-handshaking with remote",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
"reason", "local certificate is not current",
|
||||
)
|
||||
|
||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||
return
|
||||
}
|
||||
if curCrtVersion < cs.initiatingVersion {
|
||||
cm.l.WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("reason", "current cert version < pki.initiatingVersion").
|
||||
Info("Re-handshaking with remote")
|
||||
cm.l.Info("Re-handshaking with remote",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
"reason", "current cert version < pki.initiatingVersion",
|
||||
)
|
||||
|
||||
cm.intf.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], nil)
|
||||
return
|
||||
|
||||
+23
-27
@@ -7,9 +7,9 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay/overlaytest"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -46,13 +46,13 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
||||
initiatingVersion: cert.Version1,
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{version: cert.Version1},
|
||||
v1HandshakeBytes: []byte{},
|
||||
v1Credential: nil,
|
||||
}
|
||||
|
||||
lh := newTestLighthouse()
|
||||
ifce := &Interface{
|
||||
hostMap: hostMap,
|
||||
inside: &test.NoopTun{},
|
||||
inside: &overlaytest.NoopTun{},
|
||||
outside: &udp.NoopConn{},
|
||||
firewall: &Firewall{},
|
||||
lightHouse: lh,
|
||||
@@ -63,9 +63,9 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
||||
ifce.pki.cs.Store(cs)
|
||||
|
||||
// Create manager
|
||||
conf := config.NewC(l)
|
||||
punchy := NewPunchyFromConfig(l, conf)
|
||||
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||
conf := config.NewC(test.NewLogger())
|
||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||
nc.intf = ifce
|
||||
p := []byte("")
|
||||
nb := make([]byte, 12, 12)
|
||||
@@ -79,7 +79,6 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
||||
}
|
||||
hostinfo.ConnectionState = &ConnectionState{
|
||||
myCert: &dummyCert{version: cert.Version1},
|
||||
H: &noise.HandshakeState{},
|
||||
}
|
||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||
|
||||
@@ -129,13 +128,13 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
||||
initiatingVersion: cert.Version1,
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{version: cert.Version1},
|
||||
v1HandshakeBytes: []byte{},
|
||||
v1Credential: nil,
|
||||
}
|
||||
|
||||
lh := newTestLighthouse()
|
||||
ifce := &Interface{
|
||||
hostMap: hostMap,
|
||||
inside: &test.NoopTun{},
|
||||
inside: &overlaytest.NoopTun{},
|
||||
outside: &udp.NoopConn{},
|
||||
firewall: &Firewall{},
|
||||
lightHouse: lh,
|
||||
@@ -146,9 +145,9 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
||||
ifce.pki.cs.Store(cs)
|
||||
|
||||
// Create manager
|
||||
conf := config.NewC(l)
|
||||
punchy := NewPunchyFromConfig(l, conf)
|
||||
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||
conf := config.NewC(test.NewLogger())
|
||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||
nc.intf = ifce
|
||||
p := []byte("")
|
||||
nb := make([]byte, 12, 12)
|
||||
@@ -162,7 +161,6 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
||||
}
|
||||
hostinfo.ConnectionState = &ConnectionState{
|
||||
myCert: &dummyCert{version: cert.Version1},
|
||||
H: &noise.HandshakeState{},
|
||||
}
|
||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||
|
||||
@@ -214,13 +212,13 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||
initiatingVersion: cert.Version1,
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{version: cert.Version1},
|
||||
v1HandshakeBytes: []byte{},
|
||||
v1Credential: nil,
|
||||
}
|
||||
|
||||
lh := newTestLighthouse()
|
||||
ifce := &Interface{
|
||||
hostMap: hostMap,
|
||||
inside: &test.NoopTun{},
|
||||
inside: &overlaytest.NoopTun{},
|
||||
outside: &udp.NoopConn{},
|
||||
firewall: &Firewall{},
|
||||
lightHouse: lh,
|
||||
@@ -231,12 +229,12 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||
ifce.pki.cs.Store(cs)
|
||||
|
||||
// Create manager
|
||||
conf := config.NewC(l)
|
||||
conf := config.NewC(test.NewLogger())
|
||||
conf.Settings["tunnels"] = map[string]any{
|
||||
"drop_inactive": true,
|
||||
}
|
||||
punchy := NewPunchyFromConfig(l, conf)
|
||||
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||
assert.True(t, nc.dropInactive.Load())
|
||||
nc.intf = ifce
|
||||
|
||||
@@ -248,7 +246,6 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
||||
}
|
||||
hostinfo.ConnectionState = &ConnectionState{
|
||||
myCert: &dummyCert{version: cert.Version1},
|
||||
H: &noise.HandshakeState{},
|
||||
}
|
||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||
|
||||
@@ -339,15 +336,15 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
||||
cachedPeerCert, err := ncp.VerifyCertificate(now.Add(time.Second), peerCert)
|
||||
|
||||
cs := &CertState{
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{},
|
||||
v1HandshakeBytes: []byte{},
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{},
|
||||
v1Credential: nil,
|
||||
}
|
||||
|
||||
lh := newTestLighthouse()
|
||||
ifce := &Interface{
|
||||
hostMap: hostMap,
|
||||
inside: &test.NoopTun{},
|
||||
inside: &overlaytest.NoopTun{},
|
||||
outside: &udp.NoopConn{},
|
||||
firewall: &Firewall{},
|
||||
lightHouse: lh,
|
||||
@@ -360,9 +357,9 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
||||
ifce.disconnectInvalid.Store(true)
|
||||
|
||||
// Create manager
|
||||
conf := config.NewC(l)
|
||||
punchy := NewPunchyFromConfig(l, conf)
|
||||
nc := newConnectionManagerFromConfig(l, conf, hostMap, punchy)
|
||||
conf := config.NewC(test.NewLogger())
|
||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||
nc.intf = ifce
|
||||
ifce.connectionManager = nc
|
||||
|
||||
@@ -371,7 +368,6 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
||||
ConnectionState: &ConnectionState{
|
||||
myCert: &dummyCert{},
|
||||
peerCert: cachedPeerCert,
|
||||
H: &noise.HandshakeState{},
|
||||
},
|
||||
}
|
||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||
|
||||
+19
-54
@@ -1,24 +1,20 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
)
|
||||
|
||||
const ReplayWindow = 1024
|
||||
const ReplayWindow = 8192
|
||||
|
||||
type ConnectionState struct {
|
||||
eKey *NebulaCipherState
|
||||
dKey *NebulaCipherState
|
||||
H *noise.HandshakeState
|
||||
eKey noiseutil.CipherState
|
||||
dKey noiseutil.CipherState
|
||||
myCert cert.Certificate
|
||||
peerCert *cert.CachedCertificate
|
||||
initiator bool
|
||||
@@ -27,55 +23,24 @@ type ConnectionState struct {
|
||||
writeLock sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnectionState(l *logrus.Logger, cs *CertState, crt cert.Certificate, initiator bool, pattern noise.HandshakePattern) (*ConnectionState, error) {
|
||||
var dhFunc noise.DHFunc
|
||||
switch crt.Curve() {
|
||||
case cert.Curve_CURVE25519:
|
||||
dhFunc = noise.DH25519
|
||||
case cert.Curve_P256:
|
||||
if cs.pkcs11Backed {
|
||||
dhFunc = noiseutil.DHP256PKCS11
|
||||
} else {
|
||||
dhFunc = noiseutil.DHP256
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid curve: %s", crt.Curve())
|
||||
}
|
||||
|
||||
var ncs noise.CipherSuite
|
||||
if cs.cipher == "chachapoly" {
|
||||
ncs = noise.NewCipherSuite(dhFunc, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||
} else {
|
||||
ncs = noise.NewCipherSuite(dhFunc, noiseutil.CipherAESGCM, noise.HashSHA256)
|
||||
}
|
||||
|
||||
static := noise.DHKey{Private: cs.privateKey, Public: crt.PublicKey()}
|
||||
hs, err := noise.NewHandshakeState(noise.Config{
|
||||
CipherSuite: ncs,
|
||||
Random: rand.Reader,
|
||||
Pattern: pattern,
|
||||
Initiator: initiator,
|
||||
StaticKeypair: static,
|
||||
//NOTE: These should come from CertState (pki.go) when we finally implement it
|
||||
PresharedKey: []byte{},
|
||||
PresharedKeyPlacement: 0,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("NewConnectionState: %s", err)
|
||||
}
|
||||
|
||||
// The queue and ready params prevent a counter race that would happen when
|
||||
// sending stored packets and simultaneously accepting new traffic.
|
||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
||||
// completed handshake.Result. It seeds messageCounter and the replay window so
|
||||
// that the post-handshake message indices already used on the wire don't count
|
||||
// as missed traffic in the data plane.
|
||||
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
||||
ci := &ConnectionState{
|
||||
H: hs,
|
||||
initiator: initiator,
|
||||
myCert: r.MyCert,
|
||||
initiator: r.Initiator,
|
||||
peerCert: r.RemoteCert,
|
||||
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
||||
window: NewBits(ReplayWindow),
|
||||
myCert: crt,
|
||||
}
|
||||
// always start the counter from 2, as packet 1 and packet 2 are handshake packets.
|
||||
ci.messageCounter.Add(2)
|
||||
|
||||
return ci, nil
|
||||
ci.messageCounter.Add(r.MessageIndex)
|
||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
||||
ci.window.Update(nil, i)
|
||||
}
|
||||
return ci
|
||||
}
|
||||
|
||||
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/handshake"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// runTestHandshake runs a complete IX handshake between two freshly-built
|
||||
// peers and returns the initiator and responder Results. Used to produce
|
||||
// real cipher states for tests that need to exercise post-handshake glue.
|
||||
func runTestHandshake(t *testing.T) (initR, respR *handshake.Result) {
|
||||
t.Helper()
|
||||
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
|
||||
makeCreds := func(name string, networks []netip.Prefix) handshake.GetCredentialFunc {
|
||||
c, _, rawKey, _ := ct.NewTestCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
||||
name, ca.NotBefore(), ca.NotAfter(), networks, nil, nil,
|
||||
)
|
||||
priv, _, _, err := cert.UnmarshalPrivateKeyFromPEM(rawKey)
|
||||
require.NoError(t, err)
|
||||
hsBytes, err := c.MarshalForHandshakes()
|
||||
require.NoError(t, err)
|
||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||
cred := handshake.NewCredential(c, hsBytes, priv, ncs)
|
||||
return func(v cert.Version) *handshake.Credential {
|
||||
if v == cert.Version2 {
|
||||
return cred
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
verifier := func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
||||
return caPool.VerifyCertificate(time.Now(), c)
|
||||
}
|
||||
|
||||
initCreds := makeCreds("initiator", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
respCreds := makeCreds("responder", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
|
||||
initM, err := handshake.NewMachine(
|
||||
cert.Version2, initCreds, verifier,
|
||||
func() (uint32, error) { return 1000, nil },
|
||||
true, header.HandshakeIXPSK0,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
respM, err := handshake.NewMachine(
|
||||
cert.Version2, respCreds, verifier,
|
||||
func() (uint32, error) { return 2000, nil },
|
||||
false, header.HandshakeIXPSK0,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, respR, err := respM.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, respR)
|
||||
|
||||
_, initR, err = initM.ProcessPacket(nil, resp)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, initR)
|
||||
|
||||
return initR, respR
|
||||
}
|
||||
|
||||
func TestNewConnectionStateFromResult(t *testing.T) {
|
||||
initR, respR := runTestHandshake(t)
|
||||
|
||||
t.Run("initiator", func(t *testing.T) {
|
||||
ci := newConnectionStateFromResult(initR)
|
||||
assert.True(t, ci.initiator)
|
||||
assert.Equal(t, initR.MyCert, ci.myCert)
|
||||
assert.Equal(t, initR.RemoteCert, ci.peerCert)
|
||||
assert.NotNil(t, ci.eKey)
|
||||
assert.NotNil(t, ci.dKey)
|
||||
|
||||
// IX has 2 handshake messages; the next data-plane send is counter=3.
|
||||
assert.Equal(t, uint64(2), ci.messageCounter.Load(),
|
||||
"messageCounter must equal Result.MessageIndex so the next send is N+1")
|
||||
|
||||
// Both handshake counters must be marked seen so they don't appear lost.
|
||||
// Check returns false if an index has already been recorded.
|
||||
assert.False(t, ci.window.Check(nil, 1), "counter 1 must already be seen")
|
||||
assert.False(t, ci.window.Check(nil, 2), "counter 2 must already be seen")
|
||||
// Counter 3 is the next data-plane message and must NOT be pre-marked.
|
||||
assert.True(t, ci.window.Check(nil, 3), "counter 3 must not be pre-seeded")
|
||||
})
|
||||
|
||||
t.Run("responder", func(t *testing.T) {
|
||||
ci := newConnectionStateFromResult(respR)
|
||||
assert.False(t, ci.initiator)
|
||||
assert.Equal(t, respR.MyCert, ci.myCert)
|
||||
assert.Equal(t, respR.RemoteCert, ci.peerCert)
|
||||
assert.NotNil(t, ci.eKey)
|
||||
assert.NotNil(t, ci.dKey)
|
||||
assert.Equal(t, uint64(2), ci.messageCounter.Load())
|
||||
})
|
||||
}
|
||||
+80
-12
@@ -2,17 +2,33 @@ package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
"os/signal"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
)
|
||||
|
||||
type RunState int
|
||||
|
||||
const (
|
||||
StateUnknown RunState = iota
|
||||
StateReady
|
||||
StateStarted
|
||||
StateStopping
|
||||
StateStopped
|
||||
)
|
||||
|
||||
var ErrAlreadyStarted = errors.New("nebula is already started")
|
||||
var ErrAlreadyStopped = errors.New("nebula cannot be restarted")
|
||||
var ErrUnknownState = errors.New("nebula state is invalid")
|
||||
|
||||
// Every interaction here needs to take extra care to copy memory and not return or use arguments "as is" when touching
|
||||
// core. This means copying IP objects, slices, de-referencing pointers and taking the actual value, etc
|
||||
|
||||
@@ -26,8 +42,11 @@ type controlHostLister interface {
|
||||
}
|
||||
|
||||
type Control struct {
|
||||
stateLock sync.Mutex
|
||||
state RunState
|
||||
|
||||
f *Interface
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
sshStart func()
|
||||
@@ -49,10 +68,31 @@ type ControlHostInfo struct {
|
||||
CurrentRelaysThroughMe []netip.Addr `json:"currentRelaysThroughMe"`
|
||||
}
|
||||
|
||||
// Start actually runs nebula, this is a nonblocking call. To block use Control.ShutdownBlock()
|
||||
func (c *Control) Start() {
|
||||
// Start actually runs nebula, this is a nonblocking call.
|
||||
// The returned function blocks until nebula has fully stopped and returns the
|
||||
// first fatal reader error (if any). A nil error means nebula shut down
|
||||
// gracefully; a non-nil error means a reader hit an unexpected failure that
|
||||
// triggered the shutdown.
|
||||
func (c *Control) Start() (func() error, error) {
|
||||
c.stateLock.Lock()
|
||||
defer c.stateLock.Unlock()
|
||||
switch c.state {
|
||||
case StateReady:
|
||||
//yay!
|
||||
case StateStopped, StateStopping:
|
||||
return nil, ErrAlreadyStopped
|
||||
case StateStarted:
|
||||
return nil, ErrAlreadyStarted
|
||||
default:
|
||||
return nil, ErrUnknownState
|
||||
}
|
||||
|
||||
// Activate the interface
|
||||
c.f.activate()
|
||||
err := c.f.activate()
|
||||
if err != nil {
|
||||
c.state = StateStopped
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Call all the delayed funcs that waited patiently for the interface to be created.
|
||||
if c.sshStart != nil {
|
||||
@@ -71,25 +111,51 @@ func (c *Control) Start() {
|
||||
c.lighthouseStart()
|
||||
}
|
||||
|
||||
c.f.triggerShutdown = c.Stop
|
||||
|
||||
// Start reading packets.
|
||||
c.f.run()
|
||||
out, err := c.f.run()
|
||||
if err != nil {
|
||||
c.state = StateStopped
|
||||
return nil, err
|
||||
}
|
||||
c.state = StateStarted
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (c *Control) State() RunState {
|
||||
c.stateLock.Lock()
|
||||
defer c.stateLock.Unlock()
|
||||
return c.state
|
||||
}
|
||||
|
||||
func (c *Control) Context() context.Context {
|
||||
return c.ctx
|
||||
}
|
||||
|
||||
// Stop signals nebula to shutdown and close all tunnels, returns after the shutdown is complete
|
||||
// Stop is a non-blocking call that signals nebula to close all tunnels and shut down
|
||||
func (c *Control) Stop() {
|
||||
c.stateLock.Lock()
|
||||
if c.state != StateStarted {
|
||||
c.stateLock.Unlock()
|
||||
// We are stopping or stopped already
|
||||
return
|
||||
}
|
||||
|
||||
c.state = StateStopping
|
||||
c.stateLock.Unlock()
|
||||
|
||||
// Stop the handshakeManager (and other services), to prevent new tunnels from
|
||||
// being created while we're shutting them all down.
|
||||
c.cancel()
|
||||
|
||||
c.CloseAllTunnels(false)
|
||||
if err := c.f.Close(); err != nil {
|
||||
c.l.WithError(err).Error("Close interface failed")
|
||||
c.l.Error("Close interface failed", "error", err)
|
||||
}
|
||||
c.l.Info("Goodbye")
|
||||
c.stateLock.Lock()
|
||||
c.state = StateStopped
|
||||
c.stateLock.Unlock()
|
||||
}
|
||||
|
||||
// ShutdownBlock will listen for and block on term and interrupt signals, calling Control.Stop() once signalled
|
||||
@@ -100,7 +166,7 @@ func (c *Control) ShutdownBlock() {
|
||||
|
||||
rawSig := <-sigChan
|
||||
sig := rawSig.String()
|
||||
c.l.WithField("signal", sig).Info("Caught signal, shutting down")
|
||||
c.l.Info("Caught signal, shutting down", "signal", sig)
|
||||
c.Stop()
|
||||
}
|
||||
|
||||
@@ -237,8 +303,10 @@ func (c *Control) CloseAllTunnels(excludeLighthouses bool) (closed int) {
|
||||
c.f.send(header.CloseTunnel, 0, h.ConnectionState, h, []byte{}, make([]byte, 12, 12), make([]byte, mtu))
|
||||
c.f.closeTunnel(h)
|
||||
|
||||
c.l.WithField("vpnAddrs", h.vpnAddrs).WithField("udpAddr", h.remote).
|
||||
Debug("Sending close tunnel message")
|
||||
c.l.Debug("Sending close tunnel message",
|
||||
"vpnAddrs", h.vpnAddrs,
|
||||
"udpAddr", h.remote,
|
||||
)
|
||||
closed++
|
||||
}
|
||||
|
||||
|
||||
+2
-2
@@ -6,7 +6,6 @@ import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -79,10 +78,11 @@ func TestControl_GetHostInfoByVpnIp(t *testing.T) {
|
||||
}, &Interface{})
|
||||
|
||||
c := Control{
|
||||
state: StateReady,
|
||||
f: &Interface{
|
||||
hostMap: hm,
|
||||
},
|
||||
l: logrus.New(),
|
||||
l: test.NewLogger(),
|
||||
}
|
||||
|
||||
thi := c.GetHostInfoByVpnAddr(vpnIp, false)
|
||||
|
||||
+12
-60
@@ -5,8 +5,6 @@ package nebula
|
||||
import (
|
||||
"net/netip"
|
||||
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
@@ -22,7 +20,9 @@ func (c *Control) WaitForType(msgType header.MessageType, subType header.Message
|
||||
panic(err)
|
||||
}
|
||||
pipeTo.InjectUDPPacket(p)
|
||||
if h.Type == msgType && h.Subtype == subType {
|
||||
match := h.Type == msgType && h.Subtype == subType
|
||||
p.Release()
|
||||
if match {
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -38,7 +38,9 @@ func (c *Control) WaitForTypeByIndex(toIndex uint32, msgType header.MessageType,
|
||||
panic(err)
|
||||
}
|
||||
pipeTo.InjectUDPPacket(p)
|
||||
if h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType {
|
||||
match := h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType
|
||||
p.Release()
|
||||
if match {
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -90,65 +92,15 @@ func (c *Control) GetTunTxChan() <-chan []byte {
|
||||
return c.f.inside.(*overlay.TestTun).TxPackets
|
||||
}
|
||||
|
||||
// InjectUDPPacket will inject a packet into the udp side of nebula
|
||||
// InjectUDPPacket injects a packet into the udp side. We copy internally so the caller keeps ownership of p.
|
||||
// The copy comes from the freelist so steady-state alloc is zero.
|
||||
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
||||
c.f.outside.(*udp.TesterConn).Send(p)
|
||||
c.f.outside.(*udp.TesterConn).Send(p.Copy())
|
||||
}
|
||||
|
||||
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
||||
func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) {
|
||||
serialize := make([]gopacket.SerializableLayer, 0)
|
||||
var netLayer gopacket.NetworkLayer
|
||||
if toAddr.Is6() {
|
||||
if !fromAddr.Is6() {
|
||||
panic("Cant send ipv6 to ipv4")
|
||||
}
|
||||
ip := &layers.IPv6{
|
||||
Version: 6,
|
||||
NextHeader: layers.IPProtocolUDP,
|
||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||
DstIP: toAddr.Unmap().AsSlice(),
|
||||
}
|
||||
serialize = append(serialize, ip)
|
||||
netLayer = ip
|
||||
} else {
|
||||
if !fromAddr.Is4() {
|
||||
panic("Cant send ipv4 to ipv6")
|
||||
}
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Protocol: layers.IPProtocolUDP,
|
||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||
DstIP: toAddr.Unmap().AsSlice(),
|
||||
}
|
||||
serialize = append(serialize, ip)
|
||||
netLayer = ip
|
||||
}
|
||||
|
||||
udp := layers.UDP{
|
||||
SrcPort: layers.UDPPort(fromPort),
|
||||
DstPort: layers.UDPPort(toPort),
|
||||
}
|
||||
err := udp.SetNetworkLayerForChecksum(netLayer)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
buffer := gopacket.NewSerializeBuffer()
|
||||
opt := gopacket.SerializeOptions{
|
||||
ComputeChecksums: true,
|
||||
FixLengths: true,
|
||||
}
|
||||
|
||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
||||
err = gopacket.SerializeLayers(buffer, opt, serialize...)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
c.f.inside.(*overlay.TestTun).Send(buffer.Bytes())
|
||||
// InjectTunPacket pushes an IP packet onto the tun interface.
|
||||
func (c *Control) InjectTunPacket(packet []byte) {
|
||||
c.f.inside.(*overlay.TestTun).Send(packet)
|
||||
}
|
||||
|
||||
func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||
|
||||
+236
-63
@@ -1,63 +1,249 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/miekg/dns"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/config"
|
||||
)
|
||||
|
||||
// This whole thing should be rewritten to use context
|
||||
|
||||
var dnsR *dnsRecords
|
||||
var dnsServer *dns.Server
|
||||
var dnsAddr string
|
||||
|
||||
type dnsRecords struct {
|
||||
type dnsServer struct {
|
||||
sync.RWMutex
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
ctx context.Context
|
||||
dnsMap4 map[string]netip.Addr
|
||||
dnsMap6 map[string]netip.Addr
|
||||
hostMap *HostMap
|
||||
myVpnAddrsTable *bart.Lite
|
||||
|
||||
mux *dns.ServeMux
|
||||
|
||||
// enabled mirrors `lighthouse.serve_dns && lighthouse.am_lighthouse`.
|
||||
// Start, Add, and reload consult it so callers don't need to know the
|
||||
// gating rules. When it toggles off via reload, accumulated records are
|
||||
// cleared so a later re-enable starts with a fresh map populated from
|
||||
// new handshakes.
|
||||
enabled atomic.Bool
|
||||
|
||||
serverMu sync.Mutex
|
||||
server *dns.Server
|
||||
// started is closed once `server` has finished binding (or after
|
||||
// ListenAndServe returns on a bind failure). Stop waits on it before
|
||||
// calling Shutdown to avoid the miekg/dns "server not started" race
|
||||
// where a Shutdown that arrives before bind completes is silently
|
||||
// ignored, leaving the listener running forever.
|
||||
started chan struct{}
|
||||
addr string
|
||||
}
|
||||
|
||||
func newDnsRecords(l *logrus.Logger, cs *CertState, hostMap *HostMap) *dnsRecords {
|
||||
return &dnsRecords{
|
||||
// newDnsServerFromConfig builds a dnsServer, applies the initial config, and
|
||||
// registers a reload callback. The reload callback is registered before the
|
||||
// initial config is applied, so a SIGHUP can later enable, fix, or disable
|
||||
// DNS even if the initial application failed.
|
||||
//
|
||||
// The dnsServer internally gates on `lighthouse.serve_dns &&
|
||||
// lighthouse.am_lighthouse`. Start and Add are safe to call unconditionally,
|
||||
// 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
|
||||
// pointer is always non-nil, even on error.
|
||||
func newDnsServerFromConfig(ctx context.Context, l *slog.Logger, cs *CertState, hostMap *HostMap, c *config.C) (*dnsServer, error) {
|
||||
ds := &dnsServer{
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
dnsMap4: make(map[string]netip.Addr),
|
||||
dnsMap6: make(map[string]netip.Addr),
|
||||
hostMap: hostMap,
|
||||
myVpnAddrsTable: cs.myVpnAddrsTable,
|
||||
}
|
||||
ds.mux = dns.NewServeMux()
|
||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
||||
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
if err := ds.reload(c, false); err != nil {
|
||||
ds.l.Error("Failed to reload DNS responder from config", "error", err)
|
||||
}
|
||||
})
|
||||
|
||||
if err := ds.reload(c, true); err != nil {
|
||||
return ds, err
|
||||
}
|
||||
return ds, nil
|
||||
}
|
||||
|
||||
func (d *dnsRecords) Query(q uint16, data string) netip.Addr {
|
||||
// reload applies the latest config and reconciles the running state with it:
|
||||
// - enabled toggled on -> spawn a runner
|
||||
// - enabled toggled off -> stop the runner
|
||||
// - listen address changed (while running) -> restart on the new address
|
||||
// - everything else -> no-op
|
||||
//
|
||||
// On the initial call it only records configuration; Control.Start is what
|
||||
// launches the first runner via dnsStart.
|
||||
func (d *dnsServer) reload(c *config.C, initial bool) error {
|
||||
wantsDns := c.GetBool("lighthouse.serve_dns", false)
|
||||
amLighthouse := c.GetBool("lighthouse.am_lighthouse", false)
|
||||
enabled := wantsDns && amLighthouse
|
||||
newAddr := getDnsServerAddr(c)
|
||||
|
||||
d.serverMu.Lock()
|
||||
running := d.server
|
||||
runningStarted := d.started
|
||||
sameAddr := d.addr == newAddr
|
||||
d.addr = newAddr
|
||||
d.enabled.Store(enabled)
|
||||
d.serverMu.Unlock()
|
||||
|
||||
if initial {
|
||||
if wantsDns && !amLighthouse {
|
||||
d.l.Warn("DNS server refusing to run because this host is not a lighthouse.")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if !enabled {
|
||||
if running != nil {
|
||||
d.Stop()
|
||||
}
|
||||
// Drop any records that accumulated while enabled; a later re-enable
|
||||
// will repopulate from fresh handshakes.
|
||||
d.clearRecords()
|
||||
return nil
|
||||
}
|
||||
|
||||
if running == nil {
|
||||
// Was disabled (or never started); bring it up now.
|
||||
go d.Start()
|
||||
return nil
|
||||
}
|
||||
|
||||
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()
|
||||
return nil
|
||||
}
|
||||
|
||||
// shutdownServer waits for the server to finish binding (so Shutdown actually
|
||||
// stops it rather than no-oping) and then shuts it down.
|
||||
func (d *dnsServer) shutdownServer(srv *dns.Server, started chan struct{}, reason string) {
|
||||
if srv == nil {
|
||||
return
|
||||
}
|
||||
if started != nil {
|
||||
<-started
|
||||
}
|
||||
if err := srv.Shutdown(); err != nil {
|
||||
d.l.Warn("Failed to shut down the DNS responder", "reason", reason, "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Start binds and serves the DNS responder. Blocks until Stop is called or
|
||||
// the listener errors. Safe to call when DNS is disabled (returns
|
||||
// immediately). This is what Control.dnsStart points at.
|
||||
//
|
||||
// Must be invoked after the tun device is active so that lighthouse.dns.host
|
||||
// may bind to a nebula IP.
|
||||
func (d *dnsServer) Start() {
|
||||
if !d.enabled.Load() {
|
||||
return
|
||||
}
|
||||
|
||||
started := make(chan struct{})
|
||||
d.serverMu.Lock()
|
||||
if d.ctx.Err() != nil {
|
||||
d.serverMu.Unlock()
|
||||
return
|
||||
}
|
||||
addr := d.addr
|
||||
server := &dns.Server{
|
||||
Addr: addr,
|
||||
Net: "udp",
|
||||
Handler: d.mux,
|
||||
NotifyStartedFunc: func() { close(started) },
|
||||
}
|
||||
d.server = server
|
||||
d.started = started
|
||||
d.serverMu.Unlock()
|
||||
|
||||
// Per-invocation ctx watcher. Exits when Start does, so we don't leak a
|
||||
// watcher per reload-driven restart.
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
select {
|
||||
case <-d.ctx.Done():
|
||||
d.shutdownServer(server, started, "shutdown")
|
||||
case <-done:
|
||||
}
|
||||
}()
|
||||
|
||||
d.l.Info("Starting DNS responder", "dnsListener", addr)
|
||||
err := server.ListenAndServe()
|
||||
close(done)
|
||||
|
||||
// If the listener never bound (bind error) NotifyStartedFunc never fires,
|
||||
// so close started here to release any Stop caller waiting on it.
|
||||
select {
|
||||
case <-started:
|
||||
default:
|
||||
close(started)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
d.l.Warn("Failed to run the DNS responder", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Stop shuts down the active server, if any. Idempotent.
|
||||
func (d *dnsServer) Stop() {
|
||||
d.serverMu.Lock()
|
||||
srv := d.server
|
||||
started := d.started
|
||||
d.server = nil
|
||||
d.started = nil
|
||||
d.serverMu.Unlock()
|
||||
d.shutdownServer(srv, started, "stop")
|
||||
}
|
||||
|
||||
// Query returns the address for the given name and query type. The second
|
||||
// return value reports whether the name is known at all (in either A or AAAA),
|
||||
// which lets callers distinguish NODATA from NXDOMAIN.
|
||||
func (d *dnsServer) Query(q uint16, data string) (netip.Addr, bool) {
|
||||
data = strings.ToLower(data)
|
||||
d.RLock()
|
||||
defer d.RUnlock()
|
||||
addr4, haveV4 := d.dnsMap4[data]
|
||||
addr6, haveV6 := d.dnsMap6[data]
|
||||
nameExists := haveV4 || haveV6
|
||||
switch q {
|
||||
case dns.TypeA:
|
||||
if r, ok := d.dnsMap4[data]; ok {
|
||||
return r
|
||||
if haveV4 {
|
||||
return addr4, nameExists
|
||||
}
|
||||
case dns.TypeAAAA:
|
||||
if r, ok := d.dnsMap6[data]; ok {
|
||||
return r
|
||||
if haveV6 {
|
||||
return addr6, nameExists
|
||||
}
|
||||
}
|
||||
|
||||
return netip.Addr{}
|
||||
return netip.Addr{}, nameExists
|
||||
}
|
||||
|
||||
func (d *dnsRecords) QueryCert(data string) string {
|
||||
func (d *dnsServer) QueryCert(data string) string {
|
||||
if len(data) < 2 {
|
||||
return ""
|
||||
}
|
||||
ip, err := netip.ParseAddr(data[:len(data)-1])
|
||||
if err != nil {
|
||||
return ""
|
||||
@@ -80,8 +266,19 @@ func (d *dnsRecords) QueryCert(data string) string {
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// clearRecords drops all DNS records.
|
||||
func (d *dnsServer) clearRecords() {
|
||||
d.Lock()
|
||||
defer d.Unlock()
|
||||
clear(d.dnsMap4)
|
||||
clear(d.dnsMap6)
|
||||
}
|
||||
|
||||
// Add adds the first IPv4 and IPv6 address that appears in `addresses` as the record for `host`
|
||||
func (d *dnsRecords) Add(host string, addresses []netip.Addr) {
|
||||
func (d *dnsServer) Add(host string, addresses []netip.Addr) {
|
||||
if !d.enabled.Load() {
|
||||
return
|
||||
}
|
||||
host = strings.ToLower(host)
|
||||
d.Lock()
|
||||
defer d.Unlock()
|
||||
@@ -101,7 +298,7 @@ func (d *dnsRecords) Add(host string, addresses []netip.Addr) {
|
||||
}
|
||||
}
|
||||
|
||||
func (d *dnsRecords) isSelfNebulaOrLocalhost(addr string) bool {
|
||||
func (d *dnsServer) isSelfNebulaOrLocalhost(addr string) bool {
|
||||
a, _, _ := net.SplitHostPort(addr)
|
||||
b, err := netip.ParseAddr(a)
|
||||
if err != nil {
|
||||
@@ -116,13 +313,24 @@ func (d *dnsRecords) isSelfNebulaOrLocalhost(addr string) bool {
|
||||
return d.myVpnAddrsTable.Contains(b)
|
||||
}
|
||||
|
||||
func (d *dnsRecords) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||
func (d *dnsServer) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||
debugEnabled := d.l.Enabled(context.Background(), slog.LevelDebug)
|
||||
// Per RFC 2308 §2.2, a name that exists but has no record of the requested
|
||||
// type must be answered with NOERROR and an empty answer section (NODATA),
|
||||
// not NXDOMAIN (RFC 2308 §2.1), which is reserved for names that do not
|
||||
// exist at all.
|
||||
anyNameExists := false
|
||||
for _, q := range m.Question {
|
||||
switch q.Qtype {
|
||||
case dns.TypeA, dns.TypeAAAA:
|
||||
qType := dns.TypeToString[q.Qtype]
|
||||
d.l.Debugf("Query for %s %s", qType, q.Name)
|
||||
ip := d.Query(q.Qtype, q.Name)
|
||||
if debugEnabled {
|
||||
d.l.Debug("DNS query", "type", qType, "name", q.Name)
|
||||
}
|
||||
ip, nameExists := d.Query(q.Qtype, q.Name)
|
||||
if nameExists {
|
||||
anyNameExists = true
|
||||
}
|
||||
if ip.IsValid() {
|
||||
rr, err := dns.NewRR(fmt.Sprintf("%s %s %s", q.Name, qType, ip))
|
||||
if err == nil {
|
||||
@@ -134,7 +342,9 @@ func (d *dnsRecords) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||
if !d.isSelfNebulaOrLocalhost(w.RemoteAddr().String()) {
|
||||
return
|
||||
}
|
||||
d.l.Debugf("Query for TXT %s", q.Name)
|
||||
if debugEnabled {
|
||||
d.l.Debug("DNS query", "type", "TXT", "name", q.Name)
|
||||
}
|
||||
ip := d.QueryCert(q.Name)
|
||||
if ip != "" {
|
||||
rr, err := dns.NewRR(fmt.Sprintf("%s TXT %s", q.Name, ip))
|
||||
@@ -145,12 +355,12 @@ func (d *dnsRecords) parseQuery(m *dns.Msg, w dns.ResponseWriter) {
|
||||
}
|
||||
}
|
||||
|
||||
if len(m.Answer) == 0 {
|
||||
if len(m.Answer) == 0 && !anyNameExists {
|
||||
m.Rcode = dns.RcodeNameError
|
||||
}
|
||||
}
|
||||
|
||||
func (d *dnsRecords) handleDnsRequest(w dns.ResponseWriter, r *dns.Msg) {
|
||||
func (d *dnsServer) handleDnsRequest(w dns.ResponseWriter, r *dns.Msg) {
|
||||
m := new(dns.Msg)
|
||||
m.SetReply(r)
|
||||
m.Compress = false
|
||||
@@ -163,21 +373,6 @@ func (d *dnsRecords) handleDnsRequest(w dns.ResponseWriter, r *dns.Msg) {
|
||||
w.WriteMsg(m)
|
||||
}
|
||||
|
||||
func dnsMain(l *logrus.Logger, cs *CertState, hostMap *HostMap, c *config.C) func() {
|
||||
dnsR = newDnsRecords(l, cs, hostMap)
|
||||
|
||||
// attach request handler func
|
||||
dns.HandleFunc(".", dnsR.handleDnsRequest)
|
||||
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
reloadDns(l, c)
|
||||
})
|
||||
|
||||
return func() {
|
||||
startDns(l, c)
|
||||
}
|
||||
}
|
||||
|
||||
func getDnsServerAddr(c *config.C) string {
|
||||
dnsHost := strings.TrimSpace(c.GetString("lighthouse.dns.host", ""))
|
||||
// Old guidance was to provide the literal `[::]` in `lighthouse.dns.host` but that won't resolve.
|
||||
@@ -186,25 +381,3 @@ func getDnsServerAddr(c *config.C) string {
|
||||
}
|
||||
return net.JoinHostPort(dnsHost, strconv.Itoa(c.GetInt("lighthouse.dns.port", 53)))
|
||||
}
|
||||
|
||||
func startDns(l *logrus.Logger, c *config.C) {
|
||||
dnsAddr = getDnsServerAddr(c)
|
||||
dnsServer = &dns.Server{Addr: dnsAddr, Net: "udp"}
|
||||
l.WithField("dnsListener", dnsAddr).Info("Starting DNS responder")
|
||||
err := dnsServer.ListenAndServe()
|
||||
defer dnsServer.Shutdown()
|
||||
if err != nil {
|
||||
l.Errorf("Failed to start server: %s\n ", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func reloadDns(l *logrus.Logger, c *config.C) {
|
||||
if dnsAddr == getDnsServerAddr(c) {
|
||||
l.Debug("No DNS server config change detected")
|
||||
return
|
||||
}
|
||||
|
||||
l.Debug("Restarting DNS server")
|
||||
dnsServer.Shutdown()
|
||||
go startDns(l, c)
|
||||
}
|
||||
|
||||
+270
-3
@@ -1,19 +1,43 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/miekg/dns"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type stubDNSWriter struct{}
|
||||
|
||||
func (stubDNSWriter) LocalAddr() net.Addr { return &net.UDPAddr{} }
|
||||
func (stubDNSWriter) RemoteAddr() net.Addr {
|
||||
return &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 5353}
|
||||
}
|
||||
func (stubDNSWriter) Write([]byte) (int, error) { return 0, nil }
|
||||
func (stubDNSWriter) WriteMsg(*dns.Msg) error { return nil }
|
||||
func (stubDNSWriter) Close() error { return nil }
|
||||
func (stubDNSWriter) TsigStatus() error { return nil }
|
||||
func (stubDNSWriter) TsigTimersOnly(bool) {}
|
||||
func (stubDNSWriter) Hijack() {}
|
||||
|
||||
func TestParsequery(t *testing.T) {
|
||||
l := logrus.New()
|
||||
l := slog.New(slog.DiscardHandler)
|
||||
hostMap := &HostMap{}
|
||||
ds := newDnsRecords(l, &CertState{}, hostMap)
|
||||
ds := &dnsServer{
|
||||
l: l,
|
||||
dnsMap4: make(map[string]netip.Addr),
|
||||
dnsMap6: make(map[string]netip.Addr),
|
||||
hostMap: hostMap,
|
||||
}
|
||||
ds.enabled.Store(true)
|
||||
addrs := []netip.Addr{
|
||||
netip.MustParseAddr("1.2.3.4"),
|
||||
netip.MustParseAddr("1.2.3.5"),
|
||||
@@ -21,18 +45,56 @@ func TestParsequery(t *testing.T) {
|
||||
netip.MustParseAddr("fd01::25"),
|
||||
}
|
||||
ds.Add("test.com.com", addrs)
|
||||
ds.Add("v4only.com.com", []netip.Addr{netip.MustParseAddr("1.2.3.6")})
|
||||
ds.Add("v6only.com.com", []netip.Addr{netip.MustParseAddr("fd01::26")})
|
||||
|
||||
m := &dns.Msg{}
|
||||
m.SetQuestion("test.com.com", dns.TypeA)
|
||||
ds.parseQuery(m, nil)
|
||||
assert.NotNil(t, m.Answer)
|
||||
assert.Equal(t, "1.2.3.4", m.Answer[0].(*dns.A).A.String())
|
||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
||||
|
||||
m = &dns.Msg{}
|
||||
m.SetQuestion("test.com.com", dns.TypeAAAA)
|
||||
ds.parseQuery(m, nil)
|
||||
assert.NotNil(t, m.Answer)
|
||||
assert.Equal(t, "fd01::24", m.Answer[0].(*dns.AAAA).AAAA.String())
|
||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
||||
|
||||
// A known name with no record of the requested type should return NODATA
|
||||
// (NOERROR with empty answer), not NXDOMAIN.
|
||||
m = &dns.Msg{}
|
||||
m.SetQuestion("v4only.com.com", dns.TypeAAAA)
|
||||
ds.parseQuery(m, nil)
|
||||
assert.Empty(t, m.Answer)
|
||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
||||
|
||||
m = &dns.Msg{}
|
||||
m.SetQuestion("v6only.com.com", dns.TypeA)
|
||||
ds.parseQuery(m, nil)
|
||||
assert.Empty(t, m.Answer)
|
||||
assert.Equal(t, dns.RcodeSuccess, m.Rcode)
|
||||
|
||||
// An unknown name should still return NXDOMAIN.
|
||||
m = &dns.Msg{}
|
||||
m.SetQuestion("unknown.com.com", dns.TypeA)
|
||||
ds.parseQuery(m, nil)
|
||||
assert.Empty(t, m.Answer)
|
||||
assert.Equal(t, dns.RcodeNameError, m.Rcode)
|
||||
|
||||
// short lookups should not fail
|
||||
m = &dns.Msg{}
|
||||
m.Question = []dns.Question{{Name: "", Qtype: dns.TypeTXT, Qclass: dns.ClassINET}}
|
||||
ds.parseQuery(m, stubDNSWriter{})
|
||||
assert.Empty(t, m.Answer)
|
||||
assert.Equal(t, dns.RcodeNameError, m.Rcode)
|
||||
|
||||
m = &dns.Msg{}
|
||||
m.Question = []dns.Question{{Name: ".", Qtype: dns.TypeTXT, Qclass: dns.ClassINET}}
|
||||
ds.parseQuery(m, stubDNSWriter{})
|
||||
assert.Empty(t, m.Answer)
|
||||
assert.Equal(t, dns.RcodeNameError, m.Rcode)
|
||||
}
|
||||
|
||||
func Test_getDnsServerAddr(t *testing.T) {
|
||||
@@ -71,3 +133,208 @@ func Test_getDnsServerAddr(t *testing.T) {
|
||||
}
|
||||
assert.Equal(t, "[::]:1", getDnsServerAddr(c))
|
||||
}
|
||||
|
||||
func newTestDnsServer(t *testing.T) (*dnsServer, *config.C) {
|
||||
t.Helper()
|
||||
sl := slog.New(slog.DiscardHandler)
|
||||
ds := &dnsServer{
|
||||
l: sl,
|
||||
ctx: context.Background(),
|
||||
dnsMap4: make(map[string]netip.Addr),
|
||||
dnsMap6: make(map[string]netip.Addr),
|
||||
hostMap: &HostMap{},
|
||||
}
|
||||
ds.mux = dns.NewServeMux()
|
||||
ds.mux.HandleFunc(".", ds.handleDnsRequest)
|
||||
return ds, config.NewC(nil)
|
||||
}
|
||||
|
||||
func setDnsConfig(c *config.C, host string, port string, amLighthouse, serveDns bool) {
|
||||
c.Settings["lighthouse"] = map[string]any{
|
||||
"am_lighthouse": amLighthouse,
|
||||
"serve_dns": serveDns,
|
||||
"dns": map[string]any{
|
||||
"host": host,
|
||||
"port": port,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_initial_disabled(t *testing.T) {
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", "0", true, false)
|
||||
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
assert.False(t, ds.enabled.Load())
|
||||
assert.Equal(t, "127.0.0.1:0", ds.addr)
|
||||
assert.Nil(t, ds.server)
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_initial_enabled(t *testing.T) {
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
assert.True(t, ds.enabled.Load())
|
||||
assert.Equal(t, "127.0.0.1:0", ds.addr)
|
||||
// initial never starts a runner; that's Control.Start's job
|
||||
assert.Nil(t, ds.server)
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_initial_serveDnsWithoutLighthouse(t *testing.T) {
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", "0", false, true)
|
||||
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
// Wants DNS but isn't a lighthouse: gated off, no runner.
|
||||
assert.False(t, ds.enabled.Load())
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_sameAddr_noOp(t *testing.T) {
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", "0", true, true)
|
||||
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
// No server running yet, no addr change. Reload should not spawn anything.
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
assert.True(t, ds.enabled.Load())
|
||||
assert.Nil(t, ds.server)
|
||||
}
|
||||
|
||||
func TestDnsServer_StartStop_lifecycle(t *testing.T) {
|
||||
// Bind to a real (random) UDP port so we exercise the actual
|
||||
// ListenAndServe + Shutdown plumbing including the started-chan race fix.
|
||||
port := freeUDPPort(t)
|
||||
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
waitFor(t, func() bool {
|
||||
ds.serverMu.Lock()
|
||||
started := ds.started
|
||||
ds.serverMu.Unlock()
|
||||
if started == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-started:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
})
|
||||
|
||||
ds.Stop()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Start did not return after Stop")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDnsServer_Stop_beforeBind_doesNotHang(t *testing.T) {
|
||||
// Stop called immediately after Start should not deadlock even if bind
|
||||
// hasn't completed yet. This exercises the started-chan close-on-bind-fail
|
||||
// path: by binding to an obviously bad port (privileged) we get a fast
|
||||
// bind error before NotifyStartedFunc fires.
|
||||
ds, c := newTestDnsServer(t)
|
||||
// Use a port that should fail to bind (negative would be invalid, use a
|
||||
// host that won't resolve to ensure listenUDP fails quickly).
|
||||
setDnsConfig(c, "256.256.256.256", "53", true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Give Start a moment to attempt the bind and fail.
|
||||
select {
|
||||
case <-done:
|
||||
// Bind failed and Start returned; Stop should be a no-op.
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Start did not return after a bad bind")
|
||||
}
|
||||
|
||||
stopped := make(chan struct{})
|
||||
go func() {
|
||||
ds.Stop()
|
||||
close(stopped)
|
||||
}()
|
||||
select {
|
||||
case <-stopped:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Stop hung after a failed bind")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDnsServer_reload_disable_stopsRunningServer(t *testing.T) {
|
||||
port := freeUDPPort(t)
|
||||
ds, c := newTestDnsServer(t)
|
||||
setDnsConfig(c, "127.0.0.1", port, true, true)
|
||||
require.NoError(t, ds.reload(c, true))
|
||||
|
||||
startReturned := make(chan struct{})
|
||||
go func() {
|
||||
ds.Start()
|
||||
close(startReturned)
|
||||
}()
|
||||
waitForBind(t, ds)
|
||||
|
||||
// Toggle serve_dns off; reload should shut the running server down.
|
||||
setDnsConfig(c, "127.0.0.1", port, true, false)
|
||||
require.NoError(t, ds.reload(c, false))
|
||||
select {
|
||||
case <-startReturned:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Start did not return after reload disabled DNS")
|
||||
}
|
||||
assert.False(t, ds.enabled.Load())
|
||||
}
|
||||
|
||||
func freeUDPPort(t *testing.T) string {
|
||||
t.Helper()
|
||||
conn, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
port := conn.LocalAddr().(*net.UDPAddr).Port
|
||||
require.NoError(t, conn.Close())
|
||||
return strconv.Itoa(port)
|
||||
}
|
||||
|
||||
func waitForBind(t *testing.T, ds *dnsServer) {
|
||||
t.Helper()
|
||||
waitFor(t, func() bool {
|
||||
ds.serverMu.Lock()
|
||||
started := ds.started
|
||||
ds.serverMu.Unlock()
|
||||
if started == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-started:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func waitFor(t *testing.T, cond func() bool) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
if cond() {
|
||||
return
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
t.Fatal("timed out waiting for condition")
|
||||
}
|
||||
|
||||
@@ -28,6 +28,7 @@ func makeHandshakePacket(from, to netip.AddrPort, subtype header.MessageSubType,
|
||||
}
|
||||
|
||||
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Verify the responder correctly handles receiving the same msg1 multiple times
|
||||
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
||||
// and the cached response is resent.
|
||||
@@ -46,7 +47,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
||||
defer r.RenderFlow()
|
||||
|
||||
t.Log("Trigger handshake from me to them")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
|
||||
t.Log("Grab my msg1")
|
||||
msg1 := myControl.GetFromUDP(true)
|
||||
@@ -78,6 +79,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Verify that a truncated handshake packet is ignored and the real
|
||||
// packet can still complete the handshake.
|
||||
|
||||
@@ -95,7 +97,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
||||
defer r.RenderFlow()
|
||||
|
||||
t.Log("Trigger handshake")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
|
||||
t.Log("Get msg1 and deliver to responder")
|
||||
msg1 := myControl.GetFromUDP(true)
|
||||
@@ -126,6 +128,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
||||
t.Parallel()
|
||||
// A msg2 arriving with no matching pending index should be silently dropped
|
||||
// with no response sent and no state changes.
|
||||
|
||||
@@ -143,7 +146,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
||||
defer r.RenderFlow()
|
||||
|
||||
t.Log("Complete a normal handshake")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
r.RouteForAllUntilTxTun(theirControl)
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
|
||||
@@ -168,6 +171,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
||||
t.Parallel()
|
||||
// A handshake packet with an unexpected message counter should be silently
|
||||
// dropped with no side effects and no UDP response.
|
||||
|
||||
@@ -199,6 +203,7 @@ func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeUnknownSubtype(t *testing.T) {
|
||||
t.Parallel()
|
||||
// A handshake packet with an unknown subtype should be silently dropped.
|
||||
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -224,6 +229,7 @@ func TestHandshakeUnknownSubtype(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeLateResponse(t *testing.T) {
|
||||
t.Parallel()
|
||||
// After a handshake times out, a late response should be silently ignored
|
||||
// with no new tunnels created.
|
||||
|
||||
@@ -242,7 +248,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger handshake from me")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
|
||||
t.Log("Grab msg1 but don't deliver")
|
||||
msg1 := myControl.GetFromUDP(true)
|
||||
@@ -273,6 +279,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Verify that a node rejects a handshake containing its own VPN IP in the
|
||||
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
||||
|
||||
@@ -285,7 +292,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
||||
myControl.Start()
|
||||
|
||||
t.Log("Trigger handshake from me")
|
||||
myControl.InjectTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
msg1 := myControl.GetFromUDP(true)
|
||||
|
||||
t.Log("Drain any handshake retransmits before injecting")
|
||||
@@ -321,6 +328,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
||||
t.Parallel()
|
||||
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
||||
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -341,6 +349,7 @@ func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeRemoteAllowList(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Verify that a handshake from a blocked underlay IP is dropped with no
|
||||
// response and no state changes. Then verify the same packet from an
|
||||
// allowed IP succeeds.
|
||||
@@ -366,7 +375,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
||||
defer r.RenderFlow()
|
||||
|
||||
t.Log("Trigger handshake from them")
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
msg1 := theirControl.GetFromUDP(true)
|
||||
|
||||
t.Log("Rewrite the source to a blocked IP and inject")
|
||||
@@ -399,6 +408,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||
t.Parallel()
|
||||
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
||||
// remains functional and hostmap index count is stable.
|
||||
|
||||
@@ -416,7 +426,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||
defer r.RenderFlow()
|
||||
|
||||
t.Log("Complete a normal handshake via the router")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
||||
r.RouteForAllUntilTxTun(theirControl)
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
|
||||
@@ -427,7 +437,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||
originalRemote := hi.CurrentRemote
|
||||
|
||||
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam")))
|
||||
r.RouteForAllUntilTxTun(theirControl)
|
||||
|
||||
t.Log("Verify tunnel still works")
|
||||
@@ -445,6 +455,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Verify that when the wrong host responds, the cached packets are
|
||||
// transferred to the new handshake, the evil tunnel is closed, evil's
|
||||
// address is blocked, and the correct tunnel is eventually established.
|
||||
@@ -464,8 +475,8 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
||||
evilControl.Start()
|
||||
|
||||
t.Log("Send multiple packets to them (cached during handshake)")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1"))
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1")))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2")))
|
||||
|
||||
t.Log("Route until evil tunnel is closed")
|
||||
h := &header.H{}
|
||||
@@ -508,6 +519,7 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestHandshakeRelayComplete(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Verify that a relay handshake completes correctly and relay state is
|
||||
// properly maintained on all three nodes.
|
||||
|
||||
@@ -528,7 +540,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger handshake via relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
@@ -556,7 +568,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
||||
}
|
||||
|
||||
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
||||
// InjectTunUDPPacket from a V4 node to a V6 address panics in the test
|
||||
// BuildTunUDPPacket from a V4 node to a V6 address panics in the test
|
||||
// framework. The check is in handshake_manager.go handleOutbound relay
|
||||
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
||||
// address is IPv6, the relay is skipped.
|
||||
|
||||
+154
-53
@@ -11,12 +11,12 @@ import (
|
||||
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/sirupsen/logrus"
|
||||
"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/overlay"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -40,11 +40,22 @@ func BenchmarkHotPath(b *testing.B) {
|
||||
r.CancelFlowLogs()
|
||||
|
||||
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
|
||||
// Pre-build the IP packet bytes once so the bench measures the data plane,
|
||||
// not gopacket SerializeLayers overhead.
|
||||
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
|
||||
// EnableFanIn switches the router to a 0-alloc routing path. Required
|
||||
// for hot-path benchmarks; would conflict with GetFromUDP-using tests.
|
||||
r.EnableFanIn()
|
||||
|
||||
b.ResetTimer()
|
||||
|
||||
for n := 0; n < b.N; n++ {
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||
myControl.InjectTunPacket(prebuilt)
|
||||
// Release the TUN-side bytes back to the harness freelist; the bench
|
||||
// just confirms a packet arrived, the contents aren't inspected.
|
||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
||||
}
|
||||
|
||||
myControl.Stop()
|
||||
@@ -72,11 +83,15 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
||||
theirControl.Start()
|
||||
|
||||
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
|
||||
prebuilt := BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
r.EnableFanIn()
|
||||
|
||||
b.ResetTimer()
|
||||
|
||||
for n := 0; n < b.N; n++ {
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||
myControl.InjectTunPacket(prebuilt)
|
||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
||||
}
|
||||
|
||||
myControl.Stop()
|
||||
@@ -85,6 +100,7 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
||||
}
|
||||
|
||||
func TestGoodHandshake(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||
@@ -97,7 +113,7 @@ func TestGoodHandshake(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||
@@ -135,6 +151,7 @@ func TestGoodHandshake(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestGoodHandshakeNoOverlap(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, myUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
||||
@@ -170,6 +187,7 @@ func TestGoodHandshakeNoOverlap(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestWrongResponderHandshake(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
||||
@@ -189,7 +207,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
||||
evilControl.Start()
|
||||
|
||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
h := &header.H{}
|
||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
@@ -246,6 +264,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
||||
@@ -270,7 +289,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
||||
evilControl.Start()
|
||||
|
||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
h := &header.H{}
|
||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
@@ -328,6 +347,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestStage1Race(t *testing.T) {
|
||||
t.Parallel()
|
||||
// This tests ensures that two hosts handshaking with each other at the same time will allow traffic to flow
|
||||
// But will eventually collapse down to a single tunnel
|
||||
|
||||
@@ -348,8 +368,8 @@ func TestStage1Race(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake to start on both me and them")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
|
||||
t.Log("Get both stage 1 handshake packets")
|
||||
myHsForThem := myControl.GetFromUDP(true)
|
||||
@@ -408,6 +428,7 @@ func TestStage1Race(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||
@@ -425,7 +446,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
r.Log("Trigger a handshake from me to them")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
@@ -436,7 +457,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||
|
||||
myControl.InjectTunUDPPacket(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)
|
||||
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
|
||||
@@ -457,6 +478,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", nil)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||
@@ -474,7 +496,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
r.Log("Trigger a handshake from me to them")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
@@ -486,7 +508,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||
|
||||
theirControl.InjectTunUDPPacket(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)
|
||||
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
||||
@@ -508,6 +530,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
@@ -528,7 +551,7 @@ func TestRelays(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -537,6 +560,7 @@ func TestRelays(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRelaysDontCareAboutIps(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", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
||||
@@ -557,7 +581,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -566,6 +590,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestReestablishRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
@@ -586,14 +611,14 @@ func TestReestablishRelays(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||
|
||||
t.Log("Ensure packet traversal from them to me via the relay")
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
|
||||
p = r.RouteForAllUntilTxTun(myControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -608,7 +633,7 @@ func TestReestablishRelays(t *testing.T) {
|
||||
for curIndexes >= start {
|
||||
curIndexes = len(myControl.GetHostmap().Indexes)
|
||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||
myControl.InjectTunUDPPacket(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")))
|
||||
|
||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||
return router.RouteAndExit
|
||||
@@ -625,7 +650,7 @@ func TestReestablishRelays(t *testing.T) {
|
||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p = r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -660,7 +685,7 @@ func TestReestablishRelays(t *testing.T) {
|
||||
t.Log("Assert the tunnel works the other way, too")
|
||||
for {
|
||||
t.Log("RouteForAllUntilTxTun")
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
|
||||
p = r.RouteForAllUntilTxTun(myControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -697,6 +722,7 @@ func TestReestablishRelays(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestStage1RaceRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
@@ -729,8 +755,8 @@ func TestStage1RaceRelays(t *testing.T) {
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||
|
||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
|
||||
r.Log("Wait for a packet from them to me")
|
||||
p := r.RouteForAllUntilTxTun(myControl)
|
||||
@@ -744,12 +770,12 @@ func TestStage1RaceRelays(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestStage1RaceRelays2(t *testing.T) {
|
||||
t.Parallel()
|
||||
//NOTE: this is a race between me and relay resulting in a full tunnel from me to them via relay
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them ", "10.128.0.2/24", m{"relay": m{"use_relays": true}})
|
||||
l := NewTestLogger()
|
||||
|
||||
// Teach my how to get to the relay and that their can be reached via the relay
|
||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||
@@ -771,49 +797,41 @@ func TestStage1RaceRelays2(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
r.Log("Get a tunnel between me and relay")
|
||||
l.Info("Get a tunnel between me and relay")
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), myControl, relayControl, r)
|
||||
|
||||
r.Log("Get a tunnel between them and relay")
|
||||
l.Info("Get a tunnel between them and relay")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
||||
|
||||
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||
l.Info("Trigger a handshake from both them and me via relay to them and me")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
|
||||
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
||||
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
||||
|
||||
r.Log("Wait for a packet from them to me")
|
||||
l.Info("Wait for a packet from them to me; myControl")
|
||||
r.Log("Wait for a packet from them to me; myControl")
|
||||
r.RouteForAllUntilTxTun(myControl)
|
||||
l.Info("Wait for a packet from them to me; theirControl")
|
||||
r.Log("Wait for a packet from them to me; theirControl")
|
||||
r.RouteForAllUntilTxTun(theirControl)
|
||||
|
||||
r.Log("Assert the tunnel works")
|
||||
l.Info("Assert the tunnel works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
|
||||
t.Log("Wait until we remove extra tunnels")
|
||||
l.Info("Wait until we remove extra tunnels")
|
||||
l.WithFields(
|
||||
logrus.Fields{
|
||||
"myControl": len(myControl.GetHostmap().Indexes),
|
||||
"theirControl": len(theirControl.GetHostmap().Indexes),
|
||||
"relayControl": len(relayControl.GetHostmap().Indexes),
|
||||
}).Info("Waiting for hostinfos to be removed...")
|
||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||
len(myControl.GetHostmap().Indexes),
|
||||
len(theirControl.GetHostmap().Indexes),
|
||||
len(relayControl.GetHostmap().Indexes),
|
||||
)
|
||||
hostInfos := len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
||||
retries := 60
|
||||
for hostInfos > 6 && retries > 0 {
|
||||
hostInfos = len(myControl.GetHostmap().Indexes) + len(theirControl.GetHostmap().Indexes) + len(relayControl.GetHostmap().Indexes)
|
||||
l.WithFields(
|
||||
logrus.Fields{
|
||||
"myControl": len(myControl.GetHostmap().Indexes),
|
||||
"theirControl": len(theirControl.GetHostmap().Indexes),
|
||||
"relayControl": len(relayControl.GetHostmap().Indexes),
|
||||
}).Info("Waiting for hostinfos to be removed...")
|
||||
t.Logf("Waiting for hostinfos to be removed... myControl=%d theirControl=%d relayControl=%d",
|
||||
len(myControl.GetHostmap().Indexes),
|
||||
len(theirControl.GetHostmap().Indexes),
|
||||
len(relayControl.GetHostmap().Indexes),
|
||||
)
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
t.Log("Connection manager hasn't ticked yet")
|
||||
time.Sleep(time.Second)
|
||||
@@ -821,7 +839,6 @@ func TestStage1RaceRelays2(t *testing.T) {
|
||||
}
|
||||
|
||||
r.Log("Assert the tunnel works")
|
||||
l.Info("Assert the tunnel works")
|
||||
assertTunnel(t, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
||||
|
||||
myControl.Stop()
|
||||
@@ -830,6 +847,7 @@ func TestStage1RaceRelays2(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRehandshakingRelays(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||
@@ -850,7 +868,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -933,6 +951,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||
t.Parallel()
|
||||
// This test is the same as TestRehandshakingRelays but one of the terminal types is a primary swap winner
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
||||
@@ -954,7 +973,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
@@ -1037,6 +1056,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRehandshaking(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, myConfig := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.2/24", nil)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
||||
@@ -1132,6 +1152,7 @@ func TestRehandshaking(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRehandshakingLoser(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The purpose of this test is that the race loser renews their certificate and rehandshakes. The final tunnel
|
||||
// Should be the one with the new certificate
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -1230,6 +1251,7 @@ func TestRehandshakingLoser(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRaceRegression(t *testing.T) {
|
||||
t.Parallel()
|
||||
// This test forces stage 1, stage 2, stage 1 to be received by me from them
|
||||
// We had a bug where we were not finding the duplicate handshake and responding to the final stage 1 which
|
||||
// caused a cross-linked hostinfo
|
||||
@@ -1253,8 +1275,8 @@ func TestRaceRegression(t *testing.T) {
|
||||
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
||||
|
||||
t.Log("Start both handshakes")
|
||||
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
||||
|
||||
t.Log("Get both stage 1")
|
||||
myStage1ForThem := myControl.GetFromUDP(true)
|
||||
@@ -1290,6 +1312,7 @@ func TestRaceRegression(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestV2NonPrimaryWithLighthouse(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{})
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||
|
||||
@@ -1330,6 +1353,7 @@ func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestV2NonPrimaryWithOffNetLighthouse(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{})
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||
|
||||
@@ -1369,7 +1393,84 @@ func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
||||
theirControl.Stop()
|
||||
}
|
||||
|
||||
func TestLighthouseUpdateOnReload(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{})
|
||||
|
||||
// Create the lighthouse
|
||||
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh", "10.128.0.1/24", m{"lighthouse": m{"am_lighthouse": true}})
|
||||
|
||||
// Create a client with NO lighthouse configured and a long update interval.
|
||||
// The initial SendUpdate at startup will be a no-op since no lighthouses are known.
|
||||
myControl, myVpnIpNet, _, myConfig := newSimpleServer(cert.Version2, ca, caKey, "me", "10.128.0.2/24", m{
|
||||
"lighthouse": m{
|
||||
"interval": 600,
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
r := router.NewR(t, lhControl, myControl)
|
||||
defer r.RenderFlow()
|
||||
|
||||
lhControl.Start()
|
||||
myControl.Start()
|
||||
|
||||
// Drain any startup packets (there should be none meaningful)
|
||||
r.FlushAll()
|
||||
|
||||
// Verify lighthouse has no knowledge of the client
|
||||
assert.Nil(t, lhControl.QueryLighthouse(myVpnIpNet[0].Addr()))
|
||||
|
||||
// Build a new config that adds the lighthouse
|
||||
newSettings := make(m)
|
||||
for k, v := range myConfig.Settings {
|
||||
newSettings[k] = v
|
||||
}
|
||||
newSettings["static_host_map"] = m{
|
||||
lhVpnIpNet[0].Addr().String(): []any{lhUdpAddr.String()},
|
||||
}
|
||||
newSettings["lighthouse"] = m{
|
||||
"hosts": []any{lhVpnIpNet[0].Addr().String()},
|
||||
"interval": 600,
|
||||
"local_allow_list": m{
|
||||
"10.0.0.0/24": true,
|
||||
"::/0": false,
|
||||
},
|
||||
}
|
||||
newCfg, err := yaml.Marshal(newSettings)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Reload the config. The lighthouse.hosts change triggers TriggerUpdate,
|
||||
// which wakes the update worker. It calls SendUpdate, initiating a
|
||||
// handshake to the new lighthouse and caching the HostUpdateNotification.
|
||||
require.NoError(t, myConfig.ReloadConfigString(string(newCfg)))
|
||||
|
||||
// Route until the lighthouse receives the HostUpdateNotification.
|
||||
// This covers: handshake stage 1, stage 2, then the cached update.
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
r.RouteForAllUntilAfterMsgTypeTo(lhControl, header.LightHouse, 0)
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out waiting for lighthouse update after config reload")
|
||||
}
|
||||
|
||||
// Verify lighthouse now has the client's addresses
|
||||
assert.NotNil(t, lhControl.QueryLighthouse(myVpnIpNet[0].Addr()))
|
||||
|
||||
r.RenderHostmaps("Final hostmaps", lhControl, myControl)
|
||||
lhControl.Stop()
|
||||
myControl.Stop()
|
||||
}
|
||||
|
||||
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||
t.Parallel()
|
||||
unsafePrefix := "192.168.6.0/24"
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
||||
@@ -1391,7 +1492,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||
myControl.InjectTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||
@@ -1419,7 +1520,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
||||
|
||||
//reply
|
||||
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman"))
|
||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman")))
|
||||
//wait for reply
|
||||
theirControl.WaitForType(1, 0, myControl)
|
||||
theirCachedPacket := myControl.GetFromTun(true)
|
||||
|
||||
+82
-19
@@ -4,7 +4,6 @@
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net/netip"
|
||||
"os"
|
||||
@@ -12,15 +11,18 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"log/slog"
|
||||
|
||||
"dario.cat/mergo"
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.yaml.in/yaml/v3"
|
||||
@@ -132,8 +134,7 @@ func newSimpleServerWithUdpAndUnsafeNetworks(v cert.Version, caCrt cert.Certific
|
||||
"port": udpAddr.Port(),
|
||||
},
|
||||
"logging": m{
|
||||
"timestamp_format": fmt.Sprintf("%v 15:04:05.000000", name),
|
||||
"level": l.Level.String(),
|
||||
"level": testLogLevelName(),
|
||||
},
|
||||
"timers": m{
|
||||
"pending_deletion_interval": 2,
|
||||
@@ -234,8 +235,7 @@ func newServer(caCrt []cert.Certificate, certs []cert.Certificate, key []byte, o
|
||||
"port": udpAddr.Port(),
|
||||
},
|
||||
"logging": m{
|
||||
"timestamp_format": fmt.Sprintf("%v 15:04:05.000000", certs[0].Name()),
|
||||
"level": l.Level.String(),
|
||||
"level": testLogLevelName(),
|
||||
},
|
||||
"timers": m{
|
||||
"pending_deletion_interval": 2,
|
||||
@@ -294,12 +294,12 @@ func deadline(t *testing.T, seconds time.Duration) doneCb {
|
||||
|
||||
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
||||
// Send a packet from them to me
|
||||
controlB.InjectTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B"))
|
||||
controlB.InjectTunPacket(BuildTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B")))
|
||||
bPacket := r.RouteForAllUntilTxTun(controlA)
|
||||
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
||||
|
||||
// And once more from me to them
|
||||
controlA.InjectTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A"))
|
||||
controlA.InjectTunPacket(BuildTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A")))
|
||||
aPacket := r.RouteForAllUntilTxTun(controlB)
|
||||
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
||||
}
|
||||
@@ -379,24 +379,87 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
||||
return a
|
||||
}
|
||||
|
||||
func NewTestLogger() *logrus.Logger {
|
||||
l := logrus.New()
|
||||
|
||||
func NewTestLogger() *slog.Logger {
|
||||
v := os.Getenv("TEST_LOGS")
|
||||
if v == "" {
|
||||
l.SetOutput(io.Discard)
|
||||
l.SetLevel(logrus.PanicLevel)
|
||||
return l
|
||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
}
|
||||
|
||||
level := slog.LevelInfo
|
||||
switch v {
|
||||
case "2":
|
||||
l.SetLevel(logrus.DebugLevel)
|
||||
level = slog.LevelDebug
|
||||
case "3":
|
||||
l.SetLevel(logrus.TraceLevel)
|
||||
default:
|
||||
l.SetLevel(logrus.InfoLevel)
|
||||
level = logging.LevelTrace
|
||||
}
|
||||
return slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: level}))
|
||||
}
|
||||
|
||||
// testLogLevelName returns the level name string accepted by logging.ApplyConfig
|
||||
// for the current TEST_LOGS setting. Kept in sync with NewTestLogger.
|
||||
func testLogLevelName() string {
|
||||
switch os.Getenv("TEST_LOGS") {
|
||||
case "2":
|
||||
return "debug"
|
||||
case "3":
|
||||
return "trace"
|
||||
case "":
|
||||
return "info"
|
||||
}
|
||||
return "info"
|
||||
}
|
||||
|
||||
// BuildTunUDPPacket assembles an IP+UDP packet suitable for Control.InjectTunPacket.
|
||||
// Using UDP here because it's a simpler protocol.
|
||||
func BuildTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) []byte {
|
||||
serialize := make([]gopacket.SerializableLayer, 0)
|
||||
var netLayer gopacket.NetworkLayer
|
||||
if toAddr.Is6() {
|
||||
if !fromAddr.Is6() {
|
||||
panic("Cant send ipv6 to ipv4")
|
||||
}
|
||||
ip := &layers.IPv6{
|
||||
Version: 6,
|
||||
NextHeader: layers.IPProtocolUDP,
|
||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||
DstIP: toAddr.Unmap().AsSlice(),
|
||||
}
|
||||
serialize = append(serialize, ip)
|
||||
netLayer = ip
|
||||
} else {
|
||||
if !fromAddr.Is4() {
|
||||
panic("Cant send ipv4 to ipv6")
|
||||
}
|
||||
|
||||
ip := &layers.IPv4{
|
||||
Version: 4,
|
||||
TTL: 64,
|
||||
Protocol: layers.IPProtocolUDP,
|
||||
SrcIP: fromAddr.Unmap().AsSlice(),
|
||||
DstIP: toAddr.Unmap().AsSlice(),
|
||||
}
|
||||
serialize = append(serialize, ip)
|
||||
netLayer = ip
|
||||
}
|
||||
|
||||
return l
|
||||
udp := layers.UDP{
|
||||
SrcPort: layers.UDPPort(fromPort),
|
||||
DstPort: layers.UDPPort(toPort),
|
||||
}
|
||||
if err := udp.SetNetworkLayerForChecksum(netLayer); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
buffer := gopacket.NewSerializeBuffer()
|
||||
opt := gopacket.SerializeOptions{
|
||||
ComputeChecksums: true,
|
||||
FixLengths: true,
|
||||
}
|
||||
|
||||
serialize = append(serialize, &udp, gopacket.Payload(data))
|
||||
if err := gopacket.SerializeLayers(buffer, opt, serialize...); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
return buffer.Bytes()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/e2e/router"
|
||||
"go.uber.org/goleak"
|
||||
)
|
||||
|
||||
// TestNoGoroutineLeaks brings up two nebula instances, completes a tunnel,
|
||||
// stops both, and asserts no goroutines leak past the shutdown. goleak's
|
||||
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
||||
// before failing the assertion.
|
||||
//
|
||||
// Intentionally NOT t.Parallel()'d: concurrent tests would have their own
|
||||
// goroutines running and trip the assertion.
|
||||
func TestNoGoroutineLeaks(t *testing.T) {
|
||||
defer goleak.VerifyNone(t)
|
||||
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", nil)
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||
|
||||
myControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||
theirControl.InjectLightHouseAddr(myVpnIpNet[0].Addr(), myUdpAddr)
|
||||
|
||||
myControl.Start()
|
||||
theirControl.Start()
|
||||
|
||||
r := router.NewR(t, myControl, theirControl)
|
||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||
|
||||
myControl.Stop()
|
||||
theirControl.Stop()
|
||||
r.RenderFlow()
|
||||
|
||||
// Settle period: Stop() is non-blocking; the wg-driven goroutines need
|
||||
// a moment to drain. goleak retries internally too, but a short explicit
|
||||
// settle reduces flakes when the suite is busy.
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
+188
-54
@@ -13,6 +13,7 @@ import (
|
||||
"regexp"
|
||||
"sort"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -24,6 +25,19 @@ import (
|
||||
"golang.org/x/exp/maps"
|
||||
)
|
||||
|
||||
// outNatKey is the (from, to) pair used by outNat. Comparable struct, so it works as a map key without the
|
||||
// allocation cost of a string-concat key.
|
||||
type outNatKey struct {
|
||||
from, to netip.AddrPort
|
||||
}
|
||||
|
||||
// fannedPacket pairs a UDP TX packet with its source control so the router can route it after popping from
|
||||
// the fan-in channel.
|
||||
type fannedPacket struct {
|
||||
from *nebula.Control
|
||||
pkt *udp.Packet
|
||||
}
|
||||
|
||||
type R struct {
|
||||
// Simple map of the ip:port registered on a control to the control
|
||||
// Basically a router, right?
|
||||
@@ -34,12 +48,28 @@ type R struct {
|
||||
|
||||
// A last used map, if an inbound packet hit the inNat map then
|
||||
// all return packets should use the same last used inbound address for the outbound sender
|
||||
// map[from address + ":" + to address] => ip:port to rewrite in the udp packet to receiver
|
||||
outNat map[string]netip.AddrPort
|
||||
outNat map[outNatKey]netip.AddrPort
|
||||
|
||||
// A map of vpn ip to the nebula control it belongs to
|
||||
vpnControls map[netip.Addr]*nebula.Control
|
||||
|
||||
// Cached select infrastructure for RouteForAllUntilTxTun.
|
||||
// The controls map is immutable after NewR so the cases are good for the test lifetime.
|
||||
// We only rebuild if a different receiver is asked.
|
||||
selRecvCtl *nebula.Control
|
||||
selCases []reflect.SelectCase
|
||||
selCtls []*nebula.Control
|
||||
|
||||
// Optional fan-in mode for hot-path benchmarks: one forwarder goroutine per control drains UDP TX into udpFanIn,
|
||||
// so RouteForAllUntilTxTun can do a fixed 2-way native select instead of paying reflect.Select per call.
|
||||
// Off by default (would otherwise interleave with tests that use GetFromUDP directly on the same control).
|
||||
// Enabled by EnableFanIn.
|
||||
udpFanIn chan fannedPacket
|
||||
stopFanIn chan struct{}
|
||||
fanInWG sync.WaitGroup
|
||||
fanInMu sync.Mutex
|
||||
fanInOn atomic.Bool
|
||||
|
||||
ignoreFlows []ignoreFlow
|
||||
flow []flowEntry
|
||||
|
||||
@@ -119,7 +149,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
controls: make(map[netip.AddrPort]*nebula.Control),
|
||||
vpnControls: make(map[netip.Addr]*nebula.Control),
|
||||
inNat: make(map[netip.AddrPort]*nebula.Control),
|
||||
outNat: make(map[string]netip.AddrPort),
|
||||
outNat: make(map[outNatKey]netip.AddrPort),
|
||||
flow: []flowEntry{},
|
||||
ignoreFlows: []ignoreFlow{},
|
||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||
@@ -153,8 +183,10 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-clockSource.C:
|
||||
r.Lock()
|
||||
r.renderHostmaps("clock tick")
|
||||
r.renderFlow()
|
||||
r.Unlock()
|
||||
}
|
||||
}
|
||||
}()
|
||||
@@ -180,15 +212,21 @@ func (r *R) AddRoute(ip netip.Addr, port uint16, c *nebula.Control) {
|
||||
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
||||
func (r *R) RenderFlow() {
|
||||
r.cancelRender()
|
||||
r.Lock()
|
||||
defer r.Unlock()
|
||||
r.renderFlow()
|
||||
}
|
||||
|
||||
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
||||
func (r *R) CancelFlowLogs() {
|
||||
r.cancelRender()
|
||||
r.Lock()
|
||||
r.flow = nil
|
||||
r.Unlock()
|
||||
}
|
||||
|
||||
// renderFlow writes the flow log to disk. Caller must hold r.Lock. renderFlow reads r.flow / r.additionalGraphs and
|
||||
// the *packet pointers stashed inside, all of which are mutated under the same lock by routing paths.
|
||||
func (r *R) renderFlow() {
|
||||
if r.flow == nil {
|
||||
return
|
||||
@@ -434,68 +472,157 @@ func (r *R) RouteUntilTxTun(sender *nebula.Control, receiver *nebula.Control) []
|
||||
panic("No control for udp tx " + a.String())
|
||||
}
|
||||
fp := r.unlockedInjectFlow(sender, c, p, false)
|
||||
c.InjectUDPPacket(p)
|
||||
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
||||
fp.WasReceived()
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on receivers tun
|
||||
// If the router doesn't have the nebula controller for that address, we panic
|
||||
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on the receiver's tun.
|
||||
// If a control's UDP TX address can't be matched to a registered control, we panic.
|
||||
//
|
||||
// For allocation-sensitive callers (hot-path benchmarks, in particular relay
|
||||
// benches with 3+ controls), call EnableFanIn() first.
|
||||
func (r *R) RouteForAllUntilTxTun(receiver *nebula.Control) []byte {
|
||||
if r.fanInOn.Load() {
|
||||
return r.routeFanIn(receiver)
|
||||
}
|
||||
return r.routeReflect(receiver)
|
||||
}
|
||||
|
||||
// routeFanIn is the alloc-free path used when EnableFanIn is in effect.
|
||||
func (r *R) routeFanIn(receiver *nebula.Control) []byte {
|
||||
tunTx := receiver.GetTunTxChan()
|
||||
for {
|
||||
select {
|
||||
case p := <-tunTx:
|
||||
r.Lock()
|
||||
if r.flow != nil {
|
||||
np := udp.Packet{Data: make([]byte, len(p))}
|
||||
copy(np.Data, p)
|
||||
r.unlockedInjectFlow(receiver, receiver, &np, true)
|
||||
}
|
||||
r.Unlock()
|
||||
return p
|
||||
case fp := <-r.udpFanIn:
|
||||
r.routeUDP(fp.from, fp.pkt)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// routeReflect is the default reflect.Select-based path. Pays the boxing allocation per call but doesn't interfere
|
||||
// with tests that pull packets directly from controls' UDP TX channels via GetFromUDP.
|
||||
func (r *R) routeReflect(receiver *nebula.Control) []byte {
|
||||
sc, cm := r.selectCasesFor(receiver)
|
||||
for {
|
||||
x, rx, _ := reflect.Select(sc)
|
||||
if x == 0 {
|
||||
p := rx.Interface().([]byte)
|
||||
r.Lock()
|
||||
if r.flow != nil {
|
||||
np := udp.Packet{Data: make([]byte, len(p))}
|
||||
copy(np.Data, p)
|
||||
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
||||
}
|
||||
r.Unlock()
|
||||
return p
|
||||
}
|
||||
r.routeUDP(cm[x], rx.Interface().(*udp.Packet))
|
||||
}
|
||||
}
|
||||
|
||||
// EnableFanIn switches RouteForAllUntilTxTun to the alloc-free fan-in path.
|
||||
// One forwarder goroutine per registered control drains UDP TX into a shared channel that RouteForAllUntilTxTun selects
|
||||
// on alongside the receiver's TUN TX channel.
|
||||
func (r *R) EnableFanIn() {
|
||||
r.fanInMu.Lock()
|
||||
defer r.fanInMu.Unlock()
|
||||
if r.fanInOn.Load() {
|
||||
return
|
||||
}
|
||||
r.udpFanIn = make(chan fannedPacket, 32)
|
||||
r.stopFanIn = make(chan struct{})
|
||||
for _, c := range r.controls {
|
||||
r.startFanInWorker(c)
|
||||
}
|
||||
r.fanInOn.Store(true)
|
||||
r.t.Cleanup(r.stopFanInWorkers)
|
||||
}
|
||||
|
||||
// startFanInWorker spawns a goroutine that drains c's UDP TX into r.udpFanIn.
|
||||
func (r *R) startFanInWorker(c *nebula.Control) {
|
||||
r.fanInWG.Add(1)
|
||||
udpTx := c.GetUDPTxChan()
|
||||
go func() {
|
||||
defer r.fanInWG.Done()
|
||||
for {
|
||||
select {
|
||||
case <-r.stopFanIn:
|
||||
return
|
||||
case p := <-udpTx:
|
||||
select {
|
||||
case <-r.stopFanIn:
|
||||
p.Release()
|
||||
return
|
||||
case r.udpFanIn <- fannedPacket{from: c, pkt: p}:
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// stopFanInWorkers signals the fan-in goroutines to exit and waits for them.
|
||||
func (r *R) stopFanInWorkers() {
|
||||
r.fanInMu.Lock()
|
||||
wasOn := r.fanInOn.Swap(false)
|
||||
r.fanInMu.Unlock()
|
||||
if !wasOn {
|
||||
return
|
||||
}
|
||||
close(r.stopFanIn)
|
||||
r.fanInWG.Wait()
|
||||
}
|
||||
|
||||
// routeUDP forwards a UDP TX packet from the named source control to the destination control derived from p.To,
|
||||
// releasing the source packet after InjectUDPPacket has copied its bytes into a fresh pool slot.
|
||||
func (r *R) routeUDP(from *nebula.Control, p *udp.Packet) {
|
||||
r.Lock()
|
||||
defer r.Unlock()
|
||||
a := from.GetUDPAddr()
|
||||
c := r.getControl(a, p.To, p)
|
||||
if c == nil {
|
||||
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
||||
}
|
||||
fp := r.unlockedInjectFlow(from, c, p, false)
|
||||
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
||||
fp.WasReceived()
|
||||
p.Release()
|
||||
}
|
||||
|
||||
// selectCasesFor returns the SelectCase array used by routeReflect: one slot for the receiver's TUN TX channel followed
|
||||
// by one per control's UDP TX channel. Cached for the test lifetime, only rebuilt if the receiver changes.
|
||||
func (r *R) selectCasesFor(receiver *nebula.Control) ([]reflect.SelectCase, []*nebula.Control) {
|
||||
r.Lock()
|
||||
defer r.Unlock()
|
||||
if r.selRecvCtl == receiver && r.selCases != nil {
|
||||
return r.selCases, r.selCtls
|
||||
}
|
||||
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
||||
cm := make([]*nebula.Control, len(r.controls)+1)
|
||||
|
||||
i := 0
|
||||
sc[i] = reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(receiver.GetTunTxChan()),
|
||||
Send: reflect.Value{},
|
||||
}
|
||||
cm[i] = receiver
|
||||
|
||||
i++
|
||||
sc[0] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(receiver.GetTunTxChan())}
|
||||
cm[0] = receiver
|
||||
i := 1
|
||||
for _, c := range r.controls {
|
||||
sc[i] = reflect.SelectCase{
|
||||
Dir: reflect.SelectRecv,
|
||||
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||
Send: reflect.Value{},
|
||||
}
|
||||
|
||||
sc[i] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(c.GetUDPTxChan())}
|
||||
cm[i] = c
|
||||
i++
|
||||
}
|
||||
|
||||
for {
|
||||
x, rx, _ := reflect.Select(sc)
|
||||
r.Lock()
|
||||
|
||||
if x == 0 {
|
||||
// we are the tun tx, we can exit
|
||||
p := rx.Interface().([]byte)
|
||||
np := udp.Packet{Data: make([]byte, len(p))}
|
||||
copy(np.Data, p)
|
||||
|
||||
r.unlockedInjectFlow(cm[x], cm[x], &np, true)
|
||||
r.Unlock()
|
||||
return p
|
||||
|
||||
} else {
|
||||
// we are a udp tx, route and continue
|
||||
p := rx.Interface().(*udp.Packet)
|
||||
a := cm[x].GetUDPAddr()
|
||||
c := r.getControl(a, p.To, p)
|
||||
if c == nil {
|
||||
r.Unlock()
|
||||
panic(fmt.Sprintf("No control for udp tx %s", p.To))
|
||||
}
|
||||
fp := r.unlockedInjectFlow(cm[x], c, p, false)
|
||||
c.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
}
|
||||
r.Unlock()
|
||||
}
|
||||
r.selRecvCtl = receiver
|
||||
r.selCases = sc
|
||||
r.selCtls = cm
|
||||
return sc, cm
|
||||
}
|
||||
|
||||
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
||||
@@ -522,6 +649,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
||||
switch e {
|
||||
case ExitNow:
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case RouteAndExit:
|
||||
@@ -529,6 +657,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case KeepRouting:
|
||||
@@ -541,6 +670,7 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
||||
}
|
||||
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -641,6 +771,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
||||
switch e {
|
||||
case ExitNow:
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case RouteAndExit:
|
||||
@@ -648,6 +779,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
||||
receiver.InjectUDPPacket(p)
|
||||
fp.WasReceived()
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
return
|
||||
|
||||
case KeepRouting:
|
||||
@@ -659,6 +791,7 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||
}
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -702,19 +835,20 @@ func (r *R) FlushAll() {
|
||||
}
|
||||
receiver.InjectUDPPacket(p)
|
||||
r.Unlock()
|
||||
p.Release()
|
||||
}
|
||||
}
|
||||
|
||||
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
||||
// This is an internal router function, the caller must hold the lock
|
||||
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
||||
if newAddr, ok := r.outNat[fromAddr.String()+":"+toAddr.String()]; ok {
|
||||
if newAddr, ok := r.outNat[outNatKey{from: fromAddr, to: toAddr}]; ok {
|
||||
p.From = newAddr
|
||||
}
|
||||
|
||||
c, ok := r.inNat[toAddr]
|
||||
if ok {
|
||||
r.outNat[c.GetUDPAddr().String()+":"+fromAddr.String()] = toAddr
|
||||
r.outNat[outNatKey{from: c.GetUDPAddr(), to: fromAddr}] = toAddr
|
||||
return c
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
//go:build e2e_testing
|
||||
// +build e2e_testing
|
||||
|
||||
package e2e
|
||||
|
||||
import (
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"encoding/pem"
|
||||
"net"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/cert_test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/crypto/ssh"
|
||||
)
|
||||
|
||||
func TestSSHDLifecycle(t *testing.T) {
|
||||
// TestSSHDLifecycle exercises the in-process sshd through several config reloads and a Control.Stop.
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(
|
||||
cert.Version1, cert.Curve_CURVE25519,
|
||||
time.Now(), time.Now().Add(10*time.Minute),
|
||||
nil, nil, []string{},
|
||||
)
|
||||
|
||||
hostKeyPEM := generateSSHHostKey(t)
|
||||
clientSigner, clientAuthKey := generateSSHClientKey(t)
|
||||
sshdAddr := allocLoopbackPort(t)
|
||||
|
||||
overrides := m{
|
||||
"sshd": m{
|
||||
"enabled": true,
|
||||
"listen": sshdAddr,
|
||||
"host_key": hostKeyPEM,
|
||||
"authorized_users": []m{{
|
||||
"user": "tester",
|
||||
"keys": []string{clientAuthKey},
|
||||
}},
|
||||
},
|
||||
}
|
||||
control, _, _, _ := newSimpleServer(cert.Version1, ca, caKey, "sshd-test", "10.222.0.1/24", overrides)
|
||||
control.Start()
|
||||
t.Cleanup(func() { control.Stop() })
|
||||
|
||||
// sshd binds in a goroutine after Start returns; wait for it.
|
||||
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||
"sshd never started listening")
|
||||
|
||||
for i := 1; i <= 3; i++ {
|
||||
out := sshExecReload(t, sshdAddr, clientSigner)
|
||||
assert.Contains(t, out, "Reloading config", "reload cycle %d", i)
|
||||
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||
"sshd not listening after reload cycle %d", i)
|
||||
}
|
||||
|
||||
control.Stop()
|
||||
require.Eventually(t, func() bool { return !canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||
"sshd still listening after Control.Stop")
|
||||
}
|
||||
|
||||
func canDial(addr string) bool {
|
||||
c, err := net.DialTimeout("tcp", addr, 100*time.Millisecond)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
_ = c.Close()
|
||||
return true
|
||||
}
|
||||
|
||||
// allocLoopbackPort grabs an unused TCP port on 127.0.0.1, closes it, and returns the address. There
|
||||
// is a small race between releasing the port and the sshd reclaiming it; in practice the OS keeps the
|
||||
// port available long enough for the test to bind it.
|
||||
func allocLoopbackPort(t *testing.T) string {
|
||||
t.Helper()
|
||||
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
addr := l.Addr().String()
|
||||
require.NoError(t, l.Close())
|
||||
return addr
|
||||
}
|
||||
|
||||
func generateSSHHostKey(t *testing.T) string {
|
||||
t.Helper()
|
||||
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||
require.NoError(t, err)
|
||||
block, err := ssh.MarshalPrivateKey(priv, "nebula-e2e-host")
|
||||
require.NoError(t, err)
|
||||
return string(pem.EncodeToMemory(block))
|
||||
}
|
||||
|
||||
func generateSSHClientKey(t *testing.T) (ssh.Signer, string) {
|
||||
t.Helper()
|
||||
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||
require.NoError(t, err)
|
||||
signer, err := ssh.NewSignerFromKey(priv)
|
||||
require.NoError(t, err)
|
||||
auth := strings.TrimSpace(string(ssh.MarshalAuthorizedKey(signer.PublicKey())))
|
||||
return signer, auth
|
||||
}
|
||||
|
||||
func sshExecReload(t *testing.T, addr string, signer ssh.Signer) string {
|
||||
t.Helper()
|
||||
cfg := &ssh.ClientConfig{
|
||||
User: "tester",
|
||||
Auth: []ssh.AuthMethod{ssh.PublicKeys(signer)},
|
||||
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
||||
Timeout: 2 * time.Second,
|
||||
}
|
||||
client, err := ssh.Dial("tcp", addr, cfg)
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
sess, err := client.NewSession()
|
||||
require.NoError(t, err)
|
||||
defer sess.Close()
|
||||
|
||||
// reload tears the channel down before sending exit-status, so Output returns an error on the
|
||||
// channel close. The output buffer still has whatever the reload callback wrote before that.
|
||||
out, _ := sess.Output("reload")
|
||||
return string(out)
|
||||
}
|
||||
+8
-2
@@ -19,6 +19,7 @@ import (
|
||||
)
|
||||
|
||||
func TestDropInactiveTunnels(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||
// under ideal conditions
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -63,6 +64,7 @@ func TestDropInactiveTunnels(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCertUpgrade(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||
// under ideal conditions
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -157,6 +159,7 @@ func TestCertUpgrade(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCertDowngrade(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||
// under ideal conditions
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -255,6 +258,7 @@ func TestCertDowngrade(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCertMismatchCorrection(t *testing.T) {
|
||||
t.Parallel()
|
||||
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||
// under ideal conditions
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
@@ -322,6 +326,7 @@ func TestCertMismatchCorrection(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCrossStackRelaysWork(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}})
|
||||
@@ -350,14 +355,14 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
||||
theirControl.Start()
|
||||
|
||||
t.Log("Trigger a handshake from me to them via the relay")
|
||||
myControl.InjectTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me"))
|
||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
||||
|
||||
p := r.RouteForAllUntilTxTun(theirControl)
|
||||
r.Log("Assert the tunnel works")
|
||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
||||
|
||||
t.Log("reply?")
|
||||
theirControl.InjectTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them"))
|
||||
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)
|
||||
|
||||
@@ -369,6 +374,7 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCloseTunnelAuthenticated(t *testing.T) {
|
||||
t.Parallel()
|
||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.1/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "5s"}})
|
||||
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "10m"}})
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+43
-14
@@ -138,6 +138,14 @@ listen:
|
||||
# max, net.core.rmem_max and net.core.wmem_max
|
||||
#read_buffer: 10485760
|
||||
#write_buffer: 10485760
|
||||
|
||||
# On Windows only
|
||||
# When true, Nebula installs a WFP (Windows Filtering Platform) PERMIT filter scoped to UDP at the listener port.
|
||||
# WFP sits below Windows Defender Firewall, so this lets peer handshakes reach Nebula's outside socket regardless
|
||||
# of WDF's inbound rules.
|
||||
# Default true; set to false to leave WDF in charge of inbound decisions on the listener port. Not reloadable.
|
||||
#windows_bypass_wdf: true
|
||||
|
||||
# By default, Nebula replies to packets it has no tunnel for with a "recv_error" packet. This packet helps speed up reconnection
|
||||
# in the case that Nebula on either side did not shut down cleanly. This response can be abused as a way to discover if Nebula is running
|
||||
# on a host though. This option lets you configure if you want to send "recv_error" packets always, never, or only to private network remotes.
|
||||
@@ -163,17 +171,21 @@ listen:
|
||||
|
||||
punchy:
|
||||
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
||||
# This setting is reloadable.
|
||||
punch: true
|
||||
|
||||
# respond means that a node you are trying to reach will connect back out to you if your hole punching fails
|
||||
# this is extremely useful if one node is behind a difficult nat, such as a symmetric NAT
|
||||
# Default is false
|
||||
# This setting is reloadable.
|
||||
#respond: true
|
||||
|
||||
# delays a punch response for misbehaving NATs, default is 1 second.
|
||||
# This setting is reloadable.
|
||||
#delay: 1s
|
||||
|
||||
# set the delay before attempting punchy.respond. Default is 5 seconds. respond must be true to take effect.
|
||||
# This setting is reloadable.
|
||||
#respond_delay: 5s
|
||||
|
||||
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
||||
@@ -282,6 +294,24 @@ tun:
|
||||
# metric: 100
|
||||
# install: true
|
||||
|
||||
# On Windows only, sets the network category of the nebula interface. Without this, Windows often
|
||||
# leaves the network as "Unidentified" and treats it as Public, which makes the host firewall more
|
||||
# restrictive than you usually want for an overlay between trusted peers. Valid values:
|
||||
# private - treat the nebula network as a private/trusted network (default)
|
||||
# public - treat it as a public/untrusted network
|
||||
# domain - treat it as a domain-authenticated network
|
||||
# unset - leave whatever Windows decided alone
|
||||
# Not reloadable.
|
||||
#network_category: private
|
||||
|
||||
# On Windows only
|
||||
# When true, Nebula installs a WFP (Windows Filtering Platform) PERMIT filter scoped to the nebula adapter LUID.
|
||||
# WFP sits below Windows Defender Firewall, so this lets inbound traffic through regardless of WDF rules.
|
||||
# Filters are auto-removed when the adapter goes away.
|
||||
# See listen.windows_bypass_wdf for the matching control over inbound to nebula's outside UDP listener.
|
||||
# Default true; set to false to leave WDF in charge of inbound decisions on the nebula interface. Not reloadable.
|
||||
#windows_bypass_wdf: true
|
||||
|
||||
# On linux only, set to true to manage unsafe routes directly on the system route table with gateway routes instead of
|
||||
# in nebula configuration files. Default false, not reloadable.
|
||||
#use_system_route_table: false
|
||||
@@ -292,24 +322,21 @@ tun:
|
||||
|
||||
# Configure logging level
|
||||
logging:
|
||||
# panic, fatal, error, warning, info, or debug. Default is info and is reloadable.
|
||||
#NOTE: Debug mode can log remotely controlled/untrusted data which can quickly fill a disk in some
|
||||
# scenarios. Debug logging is also CPU intensive and will decrease performance overall.
|
||||
# Only enable debug logging while actively investigating an issue.
|
||||
# trace, debug, info, warn, or error. Default is info and is reloadable.
|
||||
# fatal and panic are accepted for backwards compatibility and map to error.
|
||||
#NOTE: Debug and trace modes can log remotely controlled/untrusted data which can quickly fill a disk in some
|
||||
# scenarios. Debug and trace logging are also CPU intensive and will decrease performance overall.
|
||||
# Only enable debug or trace logging while actively investigating an issue.
|
||||
level: info
|
||||
# json or text formats currently available. Default is text
|
||||
# json or text formats currently available. Default is text.
|
||||
format: text
|
||||
# Disable timestamp logging. useful when output is redirected to logging system that already adds timestamps. Default is false
|
||||
# Disable timestamp logging. Useful when output is redirected to a logging system that already adds timestamps. Default is false.
|
||||
#disable_timestamp: true
|
||||
# timestamp format is specified in Go time format, see:
|
||||
# https://golang.org/pkg/time/#pkg-constants
|
||||
# default when `format: json`: "2006-01-02T15:04:05Z07:00" (RFC3339)
|
||||
# default when `format: text`:
|
||||
# when TTY attached: seconds since beginning of execution
|
||||
# otherwise: "2006-01-02T15:04:05Z07:00" (RFC3339)
|
||||
# As an example, to log as RFC3339 with millisecond precision, set to:
|
||||
#timestamp_format: "2006-01-02T15:04:05.000Z07:00"
|
||||
# Timestamps use RFC3339Nano ("2006-01-02T15:04:05.999999999Z07:00") and are not configurable.
|
||||
|
||||
# The stats section is reloadable. A HUP may change the backend, toggle stats
|
||||
# on or off, switch the listen/host address, or pick up new DNS for the
|
||||
# configured graphite host.
|
||||
#stats:
|
||||
#type: graphite
|
||||
#prefix: nebula
|
||||
@@ -327,10 +354,12 @@ logging:
|
||||
# enables counter metrics for meta packets
|
||||
# e.g.: `messages.tx.handshake`
|
||||
# NOTE: `message.{tx,rx}.recv_error` is always emitted
|
||||
# Not reloadable.
|
||||
#message_metrics: false
|
||||
|
||||
# enables detailed counter metrics for lighthouse packets
|
||||
# e.g.: `lighthouse.rx.HostQuery`
|
||||
# Not reloadable.
|
||||
#lighthouse_metrics: false
|
||||
|
||||
# Handshake Manager Settings
|
||||
|
||||
@@ -7,9 +7,9 @@ import (
|
||||
"net"
|
||||
"os"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/service"
|
||||
)
|
||||
@@ -64,8 +64,7 @@ pki:
|
||||
return err
|
||||
}
|
||||
|
||||
logger := logrus.New()
|
||||
logger.Out = os.Stdout
|
||||
logger := logging.NewLogger(os.Stdout)
|
||||
|
||||
ctrl, err := nebula.Main(&cfg, false, "custom-app", logger, overlay.NewUserDeviceFromConfig)
|
||||
if err != nil {
|
||||
|
||||
+81
-53
@@ -1,11 +1,13 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"hash/fnv"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"slices"
|
||||
@@ -16,7 +18,6 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
@@ -67,7 +68,7 @@ type Firewall struct {
|
||||
incomingMetrics firewallMetrics
|
||||
outgoingMetrics firewallMetrics
|
||||
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type firewallMetrics struct {
|
||||
@@ -79,8 +80,8 @@ type firewallMetrics struct {
|
||||
type FirewallConntrack struct {
|
||||
sync.Mutex
|
||||
|
||||
Conns map[firewall.Packet]*conn
|
||||
TimerWheel *TimerWheel[firewall.Packet]
|
||||
Conns map[firewall.PacketKey]*conn
|
||||
TimerWheel *TimerWheel[firewall.PacketKey]
|
||||
}
|
||||
|
||||
// FirewallTable is the entry point for a rule, the evaluation order is:
|
||||
@@ -131,7 +132,7 @@ type firewallLocalCIDR struct {
|
||||
|
||||
// NewFirewall creates a new Firewall object. A TimerWheel is created for you from the provided timeouts.
|
||||
// The certificate provided should be the highest version loaded in memory.
|
||||
func NewFirewall(l *logrus.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Duration, c cert.Certificate) *Firewall {
|
||||
func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Duration, c cert.Certificate) *Firewall {
|
||||
//TODO: error on 0 duration
|
||||
var tmin, tmax time.Duration
|
||||
|
||||
@@ -165,8 +166,8 @@ func NewFirewall(l *logrus.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.D
|
||||
|
||||
return &Firewall{
|
||||
Conntrack: &FirewallConntrack{
|
||||
Conns: make(map[firewall.Packet]*conn),
|
||||
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
||||
Conns: make(map[firewall.PacketKey]*conn),
|
||||
TimerWheel: NewTimerWheel[firewall.PacketKey](tmin, tmax),
|
||||
},
|
||||
InRules: newFirewallTable(),
|
||||
OutRules: newFirewallTable(),
|
||||
@@ -191,7 +192,7 @@ func NewFirewall(l *logrus.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.D
|
||||
}
|
||||
}
|
||||
|
||||
func NewFirewallFromConfig(l *logrus.Logger, cs *CertState, c *config.C) (*Firewall, error) {
|
||||
func NewFirewallFromConfig(l *slog.Logger, cs *CertState, c *config.C) (*Firewall, error) {
|
||||
certificate := cs.getCertificate(cert.Version2)
|
||||
if certificate == nil {
|
||||
certificate = cs.getCertificate(cert.Version1)
|
||||
@@ -219,7 +220,7 @@ func NewFirewallFromConfig(l *logrus.Logger, cs *CertState, c *config.C) (*Firew
|
||||
case "drop":
|
||||
fw.InSendReject = false
|
||||
default:
|
||||
l.WithField("action", inboundAction).Warn("invalid firewall.inbound_action, defaulting to `drop`")
|
||||
l.Warn("invalid firewall.inbound_action, defaulting to `drop`", "action", inboundAction)
|
||||
fw.InSendReject = false
|
||||
}
|
||||
|
||||
@@ -230,7 +231,7 @@ func NewFirewallFromConfig(l *logrus.Logger, cs *CertState, c *config.C) (*Firew
|
||||
case "drop":
|
||||
fw.OutSendReject = false
|
||||
default:
|
||||
l.WithField("action", outboundAction).Warn("invalid firewall.outbound_action, defaulting to `drop`")
|
||||
l.Warn("invalid firewall.outbound_action, defaulting to `drop`", "action", outboundAction)
|
||||
fw.OutSendReject = false
|
||||
}
|
||||
|
||||
@@ -268,7 +269,7 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
||||
case firewall.ProtoICMP, firewall.ProtoICMPv6:
|
||||
//ICMP traffic doesn't have ports, so we always coerce to "any", even if a value is provided
|
||||
if startPort != firewall.PortAny {
|
||||
f.l.WithField("startPort", startPort).Warn("ignoring port specification for ICMP firewall rule")
|
||||
f.l.Warn("ignoring port specification for ICMP firewall rule", "startPort", startPort)
|
||||
}
|
||||
startPort = firewall.PortAny
|
||||
endPort = firewall.PortAny
|
||||
@@ -290,8 +291,9 @@ func (f *Firewall) AddRule(incoming bool, proto uint8, startPort int32, endPort
|
||||
if !incoming {
|
||||
direction = "outgoing"
|
||||
}
|
||||
f.l.WithField("firewallRule", m{"direction": direction, "proto": proto, "startPort": startPort, "endPort": endPort, "groups": groups, "host": host, "cidr": cidr, "localCidr": localCidr, "caName": caName, "caSha": caSha}).
|
||||
Info("Firewall rule added")
|
||||
f.l.Info("Firewall rule added",
|
||||
"firewallRule", m{"direction": direction, "proto": proto, "startPort": startPort, "endPort": endPort, "groups": groups, "host": host, "cidr": cidr, "localCidr": localCidr, "caName": caName, "caSha": caSha},
|
||||
)
|
||||
|
||||
return fp.addRule(f, startPort, endPort, groups, host, cidr, localCidr, caName, caSha)
|
||||
}
|
||||
@@ -314,7 +316,7 @@ func (f *Firewall) GetRuleHashes() string {
|
||||
return "SHA:" + f.GetRuleHash() + ",FNV:" + strconv.FormatUint(uint64(f.GetRuleHashFNV()), 10)
|
||||
}
|
||||
|
||||
func AddFirewallRulesFromConfig(l *logrus.Logger, inbound bool, c *config.C, fw FirewallInterface) error {
|
||||
func AddFirewallRulesFromConfig(l *slog.Logger, inbound bool, c *config.C, fw FirewallInterface) error {
|
||||
var table string
|
||||
if inbound {
|
||||
table = "firewall.inbound"
|
||||
@@ -372,7 +374,7 @@ func AddFirewallRulesFromConfig(l *logrus.Logger, inbound bool, c *config.C, fw
|
||||
startPort = firewall.PortAny
|
||||
endPort = firewall.PortAny
|
||||
if sPort != "" {
|
||||
l.WithField("port", sPort).Warn("ignoring port specification for ICMP firewall rule")
|
||||
l.Warn("ignoring port specification for ICMP firewall rule", "port", sPort)
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("%s rule #%v; proto was not understood; `%s`", table, i, r.Proto)
|
||||
@@ -396,7 +398,11 @@ func AddFirewallRulesFromConfig(l *logrus.Logger, inbound bool, c *config.C, fw
|
||||
}
|
||||
|
||||
if warning := r.sanity(); warning != nil {
|
||||
l.Warnf("%s rule #%v; %s", table, i, warning)
|
||||
l.Warn("firewall rule sanity check",
|
||||
"table", table,
|
||||
"rule", i,
|
||||
"warning", warning,
|
||||
)
|
||||
}
|
||||
|
||||
err = fw.AddRule(inbound, proto, startPort, endPort, r.Groups, r.Host, r.Cidr, r.LocalCidr, r.CAName, r.CASha)
|
||||
@@ -416,12 +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
|
||||
// returns nil if the packet should not be dropped.
|
||||
func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error {
|
||||
// Check if we spoke to this tuple, if we did then allow this packet
|
||||
if f.inConns(fp, h, caPool, localCache) {
|
||||
//
|
||||
// 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
|
||||
if h.networks == nil {
|
||||
// Simple case: Certificate has one address and no unsafe networks
|
||||
@@ -461,13 +482,13 @@ func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *
|
||||
}
|
||||
|
||||
// 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)
|
||||
return ErrNoMatchingRule
|
||||
}
|
||||
|
||||
// We always want to conntrack since it is a faster operation
|
||||
f.addConn(fp, incoming)
|
||||
f.addConn(key, fp.Protocol, incoming)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -496,9 +517,9 @@ func (f *Firewall) EmitStats() {
|
||||
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 _, ok := localCache[fp]; ok {
|
||||
if _, ok := localCache[key]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
@@ -511,7 +532,7 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
||||
f.evict(ep)
|
||||
}
|
||||
|
||||
c, ok := conntrack.Conns[fp]
|
||||
c, ok := conntrack.Conns[key]
|
||||
|
||||
if !ok {
|
||||
conntrack.Unlock()
|
||||
@@ -520,7 +541,11 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
||||
|
||||
if c.rulesVersion != f.rulesVersion {
|
||||
// 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
|
||||
if c.incoming {
|
||||
table = f.InRules
|
||||
@@ -528,32 +553,32 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
||||
|
||||
// We now know which firewall table to check against
|
||||
if !table.match(fp, c.incoming, h.ConnectionState.peerCert, caPool) {
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
h.logger(f.l).
|
||||
WithField("fwPacket", fp).
|
||||
WithField("incoming", c.incoming).
|
||||
WithField("rulesVersion", f.rulesVersion).
|
||||
WithField("oldRulesVersion", c.rulesVersion).
|
||||
Debugln("dropping old conntrack entry, does not match new ruleset")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
h.logger(f.l).Debug("dropping old conntrack entry, does not match new ruleset",
|
||||
"fwPacket", fp,
|
||||
"incoming", c.incoming,
|
||||
"rulesVersion", f.rulesVersion,
|
||||
"oldRulesVersion", c.rulesVersion,
|
||||
)
|
||||
}
|
||||
delete(conntrack.Conns, fp)
|
||||
delete(conntrack.Conns, key)
|
||||
conntrack.Unlock()
|
||||
return false
|
||||
}
|
||||
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
h.logger(f.l).
|
||||
WithField("fwPacket", fp).
|
||||
WithField("incoming", c.incoming).
|
||||
WithField("rulesVersion", f.rulesVersion).
|
||||
WithField("oldRulesVersion", c.rulesVersion).
|
||||
Debugln("keeping old conntrack entry, does match new ruleset")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
h.logger(f.l).Debug("keeping old conntrack entry, does match new ruleset",
|
||||
"fwPacket", fp,
|
||||
"incoming", c.incoming,
|
||||
"rulesVersion", f.rulesVersion,
|
||||
"oldRulesVersion", c.rulesVersion,
|
||||
)
|
||||
}
|
||||
|
||||
c.rulesVersion = f.rulesVersion
|
||||
}
|
||||
|
||||
switch fp.Protocol {
|
||||
switch key.Protocol {
|
||||
case firewall.ProtoTCP:
|
||||
c.Expires = time.Now().Add(f.TCPTimeout)
|
||||
case firewall.ProtoUDP:
|
||||
@@ -565,17 +590,17 @@ func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool,
|
||||
conntrack.Unlock()
|
||||
|
||||
if localCache != nil {
|
||||
localCache[fp] = struct{}{}
|
||||
localCache[key] = struct{}{}
|
||||
}
|
||||
|
||||
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
|
||||
c := &conn{}
|
||||
|
||||
switch fp.Protocol {
|
||||
switch protocol {
|
||||
case firewall.ProtoTCP:
|
||||
timeout = f.TCPTimeout
|
||||
case firewall.ProtoUDP:
|
||||
@@ -586,9 +611,9 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
||||
|
||||
conntrack := f.Conntrack
|
||||
conntrack.Lock()
|
||||
if _, ok := conntrack.Conns[fp]; !ok {
|
||||
if _, ok := conntrack.Conns[key]; !ok {
|
||||
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
|
||||
@@ -596,16 +621,16 @@ func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
||||
c.incoming = incoming
|
||||
c.rulesVersion = f.rulesVersion
|
||||
c.Expires = time.Now().Add(timeout)
|
||||
conntrack.Conns[fp] = c
|
||||
conntrack.Conns[key] = c
|
||||
conntrack.Unlock()
|
||||
}
|
||||
|
||||
// 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!
|
||||
func (f *Firewall) evict(p firewall.Packet) {
|
||||
func (f *Firewall) evict(key firewall.PacketKey) {
|
||||
// Are we still tracking this conn?
|
||||
conntrack := f.Conntrack
|
||||
t, ok := conntrack.Conns[p]
|
||||
t, ok := conntrack.Conns[key]
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
@@ -615,12 +640,12 @@ func (f *Firewall) evict(p firewall.Packet) {
|
||||
// Timeout is in the future, re-add the timer
|
||||
if newT > 0 {
|
||||
conntrack.TimerWheel.Advance(time.Now())
|
||||
conntrack.TimerWheel.Add(p, newT)
|
||||
conntrack.TimerWheel.Add(key, newT)
|
||||
return
|
||||
}
|
||||
|
||||
// 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 {
|
||||
@@ -935,7 +960,7 @@ type rule struct {
|
||||
CASha string
|
||||
}
|
||||
|
||||
func convertRule(l *logrus.Logger, p any, table string, i int) (rule, error) {
|
||||
func convertRule(l *slog.Logger, p any, table string, i int) (rule, error) {
|
||||
r := rule{}
|
||||
|
||||
m, ok := p.(map[string]any)
|
||||
@@ -966,7 +991,10 @@ func convertRule(l *logrus.Logger, p any, table string, i int) (rule, error) {
|
||||
return r, errors.New("group should contain a single value, an array with more than one entry was provided")
|
||||
}
|
||||
|
||||
l.Warnf("%s rule #%v; group was an array with a single value, converting to simple value", table, i)
|
||||
l.Warn("group was an array with a single value, converting to simple value",
|
||||
"table", table,
|
||||
"rule", i,
|
||||
)
|
||||
m["group"] = v[0]
|
||||
}
|
||||
|
||||
|
||||
+23
-11
@@ -1,55 +1,67 @@
|
||||
package firewall
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||
// has been seen in the conntrack table.
|
||||
type ConntrackCache map[Packet]struct{}
|
||||
// has been seen in the conntrack table. Keyed on PacketKey (dense form)
|
||||
// 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 {
|
||||
cacheV uint64
|
||||
cacheTick atomic.Uint64
|
||||
|
||||
l *slog.Logger
|
||||
cache ConntrackCache
|
||||
}
|
||||
|
||||
func NewConntrackCacheTicker(d time.Duration) *ConntrackCacheTicker {
|
||||
func NewConntrackCacheTicker(ctx context.Context, l *slog.Logger, d time.Duration) *ConntrackCacheTicker {
|
||||
if d == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
c := &ConntrackCacheTicker{
|
||||
l: l,
|
||||
cache: ConntrackCache{},
|
||||
}
|
||||
|
||||
go c.tick(d)
|
||||
go c.tick(ctx, d)
|
||||
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *ConntrackCacheTicker) tick(d time.Duration) {
|
||||
func (c *ConntrackCacheTicker) tick(ctx context.Context, d time.Duration) {
|
||||
t := time.NewTicker(d)
|
||||
defer t.Stop()
|
||||
for {
|
||||
time.Sleep(d)
|
||||
c.cacheTick.Add(1)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
c.cacheTick.Add(1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Get checks if the cache ticker has moved to the next version before returning
|
||||
// the map. If it has moved, we reset the map.
|
||||
func (c *ConntrackCacheTicker) Get(l *logrus.Logger) ConntrackCache {
|
||||
func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
||||
if c == nil {
|
||||
return nil
|
||||
}
|
||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||
c.cacheV = tick
|
||||
if ll := len(c.cache); ll > 0 {
|
||||
if l.Level == logrus.DebugLevel {
|
||||
l.WithField("len", ll).Debug("resetting conntrack cache")
|
||||
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
||||
}
|
||||
c.cache = make(ConntrackCache, ll)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
package firewall
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/test"
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
// The tests below pin the log format produced by ConntrackCacheTicker.Get
|
||||
// so changes cannot silently break what operators are grepping for. The
|
||||
// ticker's internal state (cache + cacheTick) is poked directly to avoid
|
||||
// racing a goroutine-driven tick in tests.
|
||||
|
||||
func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheTicker {
|
||||
t.Helper()
|
||||
c := &ConntrackCacheTicker{
|
||||
l: l,
|
||||
cache: make(ConntrackCache, cacheLen),
|
||||
}
|
||||
for i := 0; i < cacheLen; i++ {
|
||||
c.cache[PacketKey{LocalPort: uint16(i) + 1}] = struct{}{}
|
||||
}
|
||||
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
|
||||
return c
|
||||
}
|
||||
|
||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||
|
||||
c := newFixedTicker(t, l, 3)
|
||||
c.Get()
|
||||
|
||||
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||
}
|
||||
|
||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
||||
|
||||
c := newFixedTicker(t, l, 2)
|
||||
c.Get()
|
||||
|
||||
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||
}
|
||||
|
||||
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||
|
||||
c := newFixedTicker(t, l, 5)
|
||||
c.Get()
|
||||
|
||||
assert.Empty(t, buf.String())
|
||||
}
|
||||
|
||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||
buf := &bytes.Buffer{}
|
||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
||||
|
||||
c := newFixedTicker(t, l, 0)
|
||||
c.Get()
|
||||
|
||||
assert.Empty(t, buf.String())
|
||||
}
|
||||
@@ -19,6 +19,25 @@ const (
|
||||
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 {
|
||||
LocalAddr netip.Addr
|
||||
RemoteAddr netip.Addr
|
||||
@@ -31,6 +50,61 @@ type Packet struct {
|
||||
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 {
|
||||
return &Packet{
|
||||
LocalAddr: fp.LocalAddr,
|
||||
|
||||
+99
-109
@@ -3,13 +3,13 @@ package nebula
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"math"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
@@ -58,9 +58,8 @@ func TestNewFirewall(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFirewall_AddRule(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
|
||||
c := &dummyCert{}
|
||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c)
|
||||
@@ -177,9 +176,8 @@ func TestFirewall_AddRule(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFirewall_Drop(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||
p := firewall.Packet{
|
||||
@@ -213,50 +211,49 @@ func TestFirewall_Drop(t *testing.T) {
|
||||
cp := cert.NewCAPool()
|
||||
|
||||
// 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
|
||||
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
|
||||
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
|
||||
oldRemote := p.RemoteAddr
|
||||
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
|
||||
|
||||
// ensure signer doesn't get in the way of group checks
|
||||
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{"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
|
||||
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{"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
|
||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||
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{"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
|
||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||
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{"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) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("fd00::/7"))
|
||||
@@ -292,44 +289,44 @@ func TestFirewall_DropV6(t *testing.T) {
|
||||
cp := cert.NewCAPool()
|
||||
|
||||
// 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
|
||||
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
|
||||
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
|
||||
oldRemote := p.RemoteAddr
|
||||
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
|
||||
|
||||
// ensure signer doesn't get in the way of group checks
|
||||
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{"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
|
||||
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{"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
|
||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||
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{"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
|
||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||
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{"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) {
|
||||
@@ -485,9 +482,8 @@ func BenchmarkFirewallTable_match(b *testing.B) {
|
||||
}
|
||||
|
||||
func TestFirewall_Drop2(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||
|
||||
@@ -537,16 +533,15 @@ func TestFirewall_Drop2(t *testing.T) {
|
||||
cp := cert.NewCAPool()
|
||||
|
||||
// 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
|
||||
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) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||
|
||||
@@ -618,24 +613,23 @@ func TestFirewall_Drop3(t *testing.T) {
|
||||
cp := cert.NewCAPool()
|
||||
|
||||
// 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
|
||||
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
|
||||
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
|
||||
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.Drop(p, true, &h1, cp, nil))
|
||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil))
|
||||
}
|
||||
|
||||
func TestFirewall_Drop3V6(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("fd00::/7"))
|
||||
|
||||
@@ -667,13 +661,12 @@ func TestFirewall_Drop3V6(t *testing.T) {
|
||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||
cp := cert.NewCAPool()
|
||||
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) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||
|
||||
@@ -709,12 +702,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||
cp := cert.NewCAPool()
|
||||
|
||||
// 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
|
||||
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
|
||||
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
||||
require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
||||
|
||||
oldFw := fw
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||
@@ -723,7 +716,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||
|
||||
// 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
|
||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||
@@ -732,13 +725,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||
|
||||
// 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) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("1.1.1.1/8"))
|
||||
|
||||
@@ -778,12 +770,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 0
|
||||
p.RemotePort = 0
|
||||
// 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
|
||||
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
|
||||
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) {
|
||||
@@ -791,12 +783,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 0xabcd
|
||||
p.RemotePort = 0x1234
|
||||
// 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
|
||||
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
|
||||
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
||||
require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
|
||||
})
|
||||
})
|
||||
|
||||
@@ -808,12 +800,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 0
|
||||
p.RemotePort = 0
|
||||
// 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
|
||||
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
|
||||
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) {
|
||||
@@ -821,12 +813,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 0xabcd
|
||||
p.RemotePort = 0x1234
|
||||
// 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
|
||||
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
|
||||
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) {
|
||||
@@ -834,12 +826,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 80
|
||||
p.RemotePort = 80
|
||||
// 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
|
||||
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
|
||||
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) {
|
||||
@@ -851,12 +843,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 0
|
||||
p.RemotePort = 0
|
||||
// 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
|
||||
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
|
||||
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) {
|
||||
@@ -865,24 +857,23 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||
p.LocalPort = 0xabcd
|
||||
p.RemotePort = 0x1234
|
||||
// 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
|
||||
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
|
||||
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
|
||||
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)
|
||||
})
|
||||
})
|
||||
|
||||
}
|
||||
|
||||
func TestFirewall_DropIPSpoofing(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
myVpnNetworksTable.Insert(netip.MustParsePrefix("192.0.2.1/24"))
|
||||
|
||||
@@ -922,7 +913,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
||||
Protocol: firewall.ProtoUDP,
|
||||
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 BenchmarkLookup(b *testing.B) {
|
||||
@@ -1042,28 +1033,28 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
// Test a bad rule definition
|
||||
c := &dummyCert{}
|
||||
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil)
|
||||
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil, "aes")
|
||||
require.NoError(t, err)
|
||||
|
||||
conf := config.NewC(l)
|
||||
conf := config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": "asdf"}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound failed to parse, should be an array of rules")
|
||||
|
||||
// Test both port and code
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "code": "2"}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; only one of port or code should be provided")
|
||||
|
||||
// Test missing host, group, cidr, ca_name and ca_sha
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; at least one of host, group, cidr, local_cidr, ca_name, or ca_sha must be provided")
|
||||
|
||||
// Test code/port error
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "a", "host": "testh", "proto": "any"}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; code was not a number; `a`")
|
||||
@@ -1073,25 +1064,25 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; port was not a number; `a`")
|
||||
|
||||
// Test proto error
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "host": "testh"}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; proto was not understood; ``")
|
||||
|
||||
// Test cidr parse error
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "cidr": "testh", "proto": "any"}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; cidr did not parse; netip.ParsePrefix(\"testh\"): no '/'")
|
||||
|
||||
// Test local_cidr parse error
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"code": "1", "local_cidr": "testh", "proto": "any"}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.outbound rule #0; local_cidr did not parse; netip.ParsePrefix(\"testh\"): no '/'")
|
||||
|
||||
// Test both group and groups
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "group": "a", "groups": []string{"b", "c"}}}}
|
||||
_, err = NewFirewallFromConfig(l, cs, conf)
|
||||
require.EqualError(t, err, "firewall.inbound rule #0; only one of group or groups should be defined, both provided")
|
||||
@@ -1100,35 +1091,35 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
||||
func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
// Test adding tcp rule
|
||||
conf := config.NewC(l)
|
||||
conf := config.NewC(test.NewLogger())
|
||||
mf := &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "tcp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoTCP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding udp rule
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "udp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoUDP, startPort: 1, endPort: 1, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding icmp rule
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"port": "1", "proto": "icmp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding icmp rule no port
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"outbound": []any{map[string]any{"proto": "icmp", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, false, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: false, proto: firewall.ProtoICMP, startPort: firewall.PortAny, endPort: firewall.PortAny, groups: nil, host: "a", ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding any rule
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "host": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
@@ -1136,14 +1127,14 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||
|
||||
// Test adding rule with cidr
|
||||
cidr := netip.MustParsePrefix("10.0.0.0/8")
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": cidr.String()}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: cidr.String(), localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding rule with local_cidr
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": cidr.String()}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
@@ -1151,82 +1142,82 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||
|
||||
// Test adding rule with cidr ipv6
|
||||
cidr6 := netip.MustParsePrefix("fd00::/8")
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": cidr6.String()}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: cidr6.String(), localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding rule with any cidr
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": "any"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "any", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test adding rule with junk cidr
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "cidr": "junk/junk"}}}
|
||||
require.EqualError(t, AddFirewallRulesFromConfig(l, true, conf, mf), "firewall.inbound rule #0; cidr did not parse; netip.ParsePrefix(\"junk/junk\"): ParseAddr(\"junk\"): unable to parse IP")
|
||||
|
||||
// Test adding rule with local_cidr ipv6
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": cidr6.String()}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: cidr6.String()}, mf.lastCall)
|
||||
|
||||
// Test adding rule with any local_cidr
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": "any"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, localIp: "any"}, mf.lastCall)
|
||||
|
||||
// Test adding rule with junk local_cidr
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "local_cidr": "junk/junk"}}}
|
||||
require.EqualError(t, AddFirewallRulesFromConfig(l, true, conf, mf), "firewall.inbound rule #0; local_cidr did not parse; netip.ParsePrefix(\"junk/junk\"): ParseAddr(\"junk\"): unable to parse IP")
|
||||
|
||||
// Test adding rule with ca_sha
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "ca_sha": "12312313123"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: "", caSha: "12312313123"}, mf.lastCall)
|
||||
|
||||
// Test adding rule with ca_name
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "ca_name": "root01"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: nil, ip: "", localIp: "", caName: "root01"}, mf.lastCall)
|
||||
|
||||
// Test single group
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "group": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a"}, ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test single groups
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "groups": "a"}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a"}, ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test multiple AND groups
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "groups": []string{"a", "b"}}}}
|
||||
require.NoError(t, AddFirewallRulesFromConfig(l, true, conf, mf))
|
||||
assert.Equal(t, addRuleCall{incoming: true, proto: firewall.ProtoAny, startPort: 1, endPort: 1, groups: []string{"a", "b"}, ip: "", localIp: ""}, mf.lastCall)
|
||||
|
||||
// Test Add error
|
||||
conf = config.NewC(l)
|
||||
conf = config.NewC(test.NewLogger())
|
||||
mf = &mockFirewall{}
|
||||
mf.nextCallReturn = errors.New("test error")
|
||||
conf.Settings["firewall"] = map[string]any{"inbound": []any{map[string]any{"port": "1", "proto": "any", "host": "a"}}}
|
||||
@@ -1234,9 +1225,8 @@ func TestAddFirewallRulesFromConfig(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFirewall_convertRule(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
|
||||
// Ensure group array of 1 is converted and a warning is printed
|
||||
c := map[string]any{
|
||||
@@ -1244,7 +1234,9 @@ func TestFirewall_convertRule(t *testing.T) {
|
||||
}
|
||||
|
||||
r, err := convertRule(l, c, "test", 1)
|
||||
assert.Contains(t, ob.String(), "test rule #1; group was an array with a single value, converting to simple value")
|
||||
assert.Contains(t, ob.String(), "group was an array with a single value, converting to simple value")
|
||||
assert.Contains(t, ob.String(), "table=test")
|
||||
assert.Contains(t, ob.String(), "rule=1")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"group1"}, r.Groups)
|
||||
|
||||
@@ -1270,9 +1262,8 @@ func TestFirewall_convertRule(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFirewall_convertRuleSanity(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
|
||||
noWarningPlease := []map[string]any{
|
||||
{"group": "group1"},
|
||||
@@ -1336,7 +1327,7 @@ func (c *testcase) Test(t *testing.T, fw *Firewall) {
|
||||
t.Helper()
|
||||
cp := cert.NewCAPool()
|
||||
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 {
|
||||
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
||||
} else {
|
||||
@@ -1386,7 +1377,7 @@ type testsetup struct {
|
||||
fw *Firewall
|
||||
}
|
||||
|
||||
func newSetup(t *testing.T, l *logrus.Logger, myPrefixes ...netip.Prefix) testsetup {
|
||||
func newSetup(t *testing.T, l *slog.Logger, myPrefixes ...netip.Prefix) testsetup {
|
||||
c := dummyCert{
|
||||
name: "me",
|
||||
networks: myPrefixes,
|
||||
@@ -1397,7 +1388,7 @@ func newSetup(t *testing.T, l *logrus.Logger, myPrefixes ...netip.Prefix) testse
|
||||
return newSetupFromCert(t, l, c)
|
||||
}
|
||||
|
||||
func newSetupFromCert(t *testing.T, l *logrus.Logger, c dummyCert) testsetup {
|
||||
func newSetupFromCert(t *testing.T, l *slog.Logger, c dummyCert) testsetup {
|
||||
myVpnNetworksTable := new(bart.Lite)
|
||||
for _, prefix := range c.Networks() {
|
||||
myVpnNetworksTable.Insert(prefix)
|
||||
@@ -1414,9 +1405,8 @@ func newSetupFromCert(t *testing.T, l *logrus.Logger, c dummyCert) testsetup {
|
||||
|
||||
func TestFirewall_Drop_EnforceIPMatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
l := test.NewLogger()
|
||||
ob := &bytes.Buffer{}
|
||||
l.SetOutput(ob)
|
||||
l := test.NewLoggerWithOutput(ob)
|
||||
|
||||
myPrefix := netip.MustParsePrefix("1.1.1.1/8")
|
||||
// for now, it's okay that these are all "incoming", the logic this test tries to check doesn't care about in/out
|
||||
@@ -1529,6 +1519,6 @@ func (mf *mockFirewall) AddRule(incoming bool, proto uint8, startPort int32, end
|
||||
|
||||
func resetConntrack(fw *Firewall) {
|
||||
fw.Conntrack.Lock()
|
||||
fw.Conntrack.Conns = map[firewall.Packet]*conn{}
|
||||
fw.Conntrack.Conns = map[firewall.PacketKey]*conn{}
|
||||
fw.Conntrack.Unlock()
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
module github.com/slackhq/nebula
|
||||
|
||||
go 1.25
|
||||
go 1.25.0
|
||||
|
||||
require (
|
||||
dario.cat/mergo v1.0.2
|
||||
@@ -9,7 +9,7 @@ require (
|
||||
github.com/armon/go-radix v1.0.0
|
||||
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
||||
github.com/flynn/noise v1.1.0
|
||||
github.com/gaissmai/bart v0.26.0
|
||||
github.com/gaissmai/bart v0.26.1
|
||||
github.com/gogo/protobuf v1.3.2
|
||||
github.com/google/gopacket v1.1.19
|
||||
github.com/kardianos/service v1.2.4
|
||||
@@ -18,21 +18,21 @@ require (
|
||||
github.com/nbrownus/go-metrics-prometheus v0.0.0-20210712211119-974a6260965f
|
||||
github.com/prometheus/client_golang v1.23.2
|
||||
github.com/rcrowley/go-metrics v0.0.0-20201227073835-cf1acfcdf475
|
||||
github.com/sirupsen/logrus v1.9.4
|
||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
|
||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||
github.com/stretchr/testify v1.11.1
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
go.uber.org/goleak v1.3.0
|
||||
go.yaml.in/yaml/v3 v3.0.4
|
||||
golang.org/x/crypto v0.47.0
|
||||
golang.org/x/crypto v0.50.0
|
||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||
golang.org/x/net v0.49.0
|
||||
golang.org/x/sync v0.19.0
|
||||
golang.org/x/sys v0.40.0
|
||||
golang.org/x/term v0.39.0
|
||||
golang.org/x/net v0.53.0
|
||||
golang.org/x/sync v0.20.0
|
||||
golang.org/x/sys v0.43.0
|
||||
golang.org/x/term v0.42.0
|
||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b
|
||||
golang.zx2c4.com/wireguard/windows v0.5.3
|
||||
golang.zx2c4.com/wireguard/windows v0.6.1
|
||||
google.golang.org/protobuf v1.36.11
|
||||
gopkg.in/yaml.v3 v3.0.1
|
||||
gvisor.dev/gvisor v0.0.0-20240423190808-9d7a357edefe
|
||||
@@ -43,6 +43,7 @@ require (
|
||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||
github.com/davecgh/go-spew v1.1.1 // 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/pmezard/go-difflib v1.0.0 // 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/vishvananda/netns v0.0.5 // indirect
|
||||
go.yaml.in/yaml/v2 v2.4.2 // indirect
|
||||
golang.org/x/mod v0.31.0 // indirect
|
||||
golang.org/x/mod v0.34.0 // indirect
|
||||
golang.org/x/time v0.5.0 // indirect
|
||||
golang.org/x/tools v0.40.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/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
||||
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
||||
github.com/gaissmai/bart v0.26.0 h1:xOZ57E9hJLBiQaSyeZa9wgWhGuzfGACgqp4BE77OkO0=
|
||||
github.com/gaissmai/bart v0.26.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
||||
github.com/gaissmai/bart v0.26.1 h1:+w4rnLGNlA2GDVn382Tfe3jOsK5vOr5n4KmigJ9lbTo=
|
||||
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.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||
@@ -60,6 +60,8 @@ 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/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8=
|
||||
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/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU=
|
||||
github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
|
||||
@@ -133,8 +135,6 @@ github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncj
|
||||
github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo=
|
||||
github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE=
|
||||
github.com/sirupsen/logrus v1.6.0/go.mod h1:7uNnSEd1DgxDLC74fIahvMZmmYsHGZGEOFrfsX/uA88=
|
||||
github.com/sirupsen/logrus v1.9.4 h1:TsZE7l11zFCLZnZ+teH4Umoq5BhEIfIzfRDZ1Uzql2w=
|
||||
github.com/sirupsen/logrus v1.9.4/go.mod h1:ftWc9WdOfJ0a92nsE2jF5u5ZwH8Bv2zdeOC42RjbV2g=
|
||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e h1:MRM5ITcdelLK2j1vwZ3Je0FKVCfqOLp5zO6trqMLYs0=
|
||||
github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e/go.mod h1:XV66xRDqSt+GTGFMVlhk3ULuV0y9ZmzeVGR4mloJI3M=
|
||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6 h1:pnnLyeX7o/5aX8qUQ69P/mLojDqwda8hFOCBTmP/6hw=
|
||||
@@ -164,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-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
|
||||
golang.org/x/crypto v0.0.0-20210322153248-0c34fe9e7dc2/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4=
|
||||
golang.org/x/crypto v0.47.0 h1:V6e3FRj+n4dbpw86FJ8Fv7XVOql7TEwpHapKoMJ/GO8=
|
||||
golang.org/x/crypto v0.47.0/go.mod h1:ff3Y9VzzKbwSSEzWqJsJVBnWmRwRSHt/6Op5n9bQc4A=
|
||||
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
|
||||
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/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg=
|
||||
golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA=
|
||||
golang.org/x/mod v0.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
|
||||
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
|
||||
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
||||
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
|
||||
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
|
||||
@@ -184,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-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.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
|
||||
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
|
||||
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
||||
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/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=
|
||||
@@ -193,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-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.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
|
||||
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
|
||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
@@ -210,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.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ=
|
||||
golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
|
||||
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
|
||||
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.39.0 h1:RclSuaJf32jOqZz74CkPA9qFuVTX7vhLlpfj/IGWlqY=
|
||||
golang.org/x/term v0.39.0/go.mod h1:yxzUCTP/U+FzoxfdKmLaA0RV1WgE0VY7hXBwKtY/4ww=
|
||||
golang.org/x/term v0.42.0 h1:UiKe+zDFmJobeJ5ggPwOshJIVt6/Ft0rcfrXZDLWAWY=
|
||||
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.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
|
||||
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||
@@ -225,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-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.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
|
||||
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
|
||||
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
|
||||
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
|
||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||
@@ -235,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/wireguard v0.0.0-20230325221338-052af4a8072b h1:J1CaxgLerRR5lgx3wnr6L04cJFbWoceSK9JWBdglINo=
|
||||
golang.zx2c4.com/wireguard v0.0.0-20230325221338-052af4a8072b/go.mod h1:tqur9LnfstdR9ep2LaJT4lFUl0EjlHtge+gAjmsHUG4=
|
||||
golang.zx2c4.com/wireguard/windows v0.5.3 h1:On6j2Rpn3OEMXqBq00QEDC7bWSZrPIHKIus8eIuExIE=
|
||||
golang.zx2c4.com/wireguard/windows v0.5.3/go.mod h1:9TEe8TJmtwyQebdFwAkEWOPr3prrtqm+REGFifP60hI=
|
||||
golang.zx2c4.com/wireguard/windows v0.6.1 h1:XMaKojH1Hs/raMrmnir4n35nTvzvWj7NmSYzHn2F4qU=
|
||||
golang.zx2c4.com/wireguard/windows v0.6.1/go.mod h1:04aqInu5GYuTFvMuDw/rKBAF7mHrltW/3rekpfbbZDM=
|
||||
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
|
||||
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
|
||||
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
)
|
||||
|
||||
// Credential holds everything needed to participate in a handshake
|
||||
// at a given cert version. Version and Curve are read from Cert; the public
|
||||
// half of the static keypair likewise comes from Cert.PublicKey().
|
||||
type Credential struct {
|
||||
Cert cert.Certificate // the certificate
|
||||
Bytes []byte // pre-marshaled certificate bytes
|
||||
privateKey []byte // static private key (public half lives in Cert)
|
||||
cipherSuite noise.CipherSuite // pre-built cipher suite (DH + cipher + hash)
|
||||
}
|
||||
|
||||
// NewCredential creates a Credential with all material needed for handshake
|
||||
// participation. The cipherSuite should be pre-built by the caller with the
|
||||
// appropriate DH function, cipher, and hash.
|
||||
func NewCredential(
|
||||
c cert.Certificate,
|
||||
hsBytes []byte,
|
||||
privateKey []byte,
|
||||
cipherSuite noise.CipherSuite,
|
||||
) *Credential {
|
||||
return &Credential{
|
||||
Cert: c,
|
||||
Bytes: hsBytes,
|
||||
privateKey: privateKey,
|
||||
cipherSuite: cipherSuite,
|
||||
}
|
||||
}
|
||||
|
||||
// buildHandshakeState creates a noise.HandshakeState from this credential.
|
||||
func (hc *Credential) buildHandshakeState(initiator bool, pattern noise.HandshakePattern) (*noise.HandshakeState, error) {
|
||||
return noise.NewHandshakeState(noise.Config{
|
||||
CipherSuite: hc.cipherSuite,
|
||||
Random: rand.Reader,
|
||||
Pattern: pattern,
|
||||
Initiator: initiator,
|
||||
StaticKeypair: noise.DHKey{Private: hc.privateKey, Public: hc.Cert.PublicKey()},
|
||||
PresharedKey: []byte{},
|
||||
PresharedKeyPlacement: 0,
|
||||
})
|
||||
}
|
||||
|
||||
// GetCredentialFunc returns the handshake credential for the given version,
|
||||
// or nil if that version is not available.
|
||||
//
|
||||
// Implementations must return credentials drawn from a snapshot stable for
|
||||
// the lifetime of any single Machine. The Machine may call this multiple
|
||||
// times during a handshake (e.g. when negotiating to the peer's version)
|
||||
// and assumes the underlying static keypair is consistent across calls.
|
||||
type GetCredentialFunc func(v cert.Version) *Credential
|
||||
@@ -0,0 +1,21 @@
|
||||
package handshake
|
||||
|
||||
import "errors"
|
||||
|
||||
var (
|
||||
ErrInitiateOnResponder = errors.New("initiate called on responder")
|
||||
ErrInitiateAlreadyCalled = errors.New("initiate already called")
|
||||
ErrInitiateNotCalled = errors.New("initiate must be called before ProcessPacket for initiators")
|
||||
ErrPacketTooShort = errors.New("packet too short")
|
||||
ErrPublicKeyMismatch = errors.New("public key mismatch between certificate and handshake")
|
||||
ErrIncompleteHandshake = errors.New("handshake completed without receiving required content")
|
||||
ErrMachineFailed = errors.New("handshake machine has failed")
|
||||
ErrUnknownSubtype = errors.New("unknown handshake subtype")
|
||||
ErrMissingContent = errors.New("expected handshake content but message was empty")
|
||||
ErrUnexpectedContent = errors.New("received unexpected handshake content")
|
||||
ErrIndexAllocation = errors.New("failed to allocate local index")
|
||||
ErrNoCredential = errors.New("no handshake credential available for cert version")
|
||||
ErrAsymmetricCipherKeys = errors.New("noise produced only one cipher key")
|
||||
ErrMultiMessageUnsupported = errors.New("multi-message handshake patterns are not yet supported by the manager")
|
||||
ErrSubtypeMismatch = errors.New("packet subtype does not match handshake machine subtype")
|
||||
)
|
||||
@@ -0,0 +1,29 @@
|
||||
// This file documents the wire format the nebula handshake speaks. It is
|
||||
// not run through protoc; the encoder/decoder in payload.go is hand-written
|
||||
// against this shape directly to keep the parser narrow and panic-free.
|
||||
//
|
||||
// Any change to the wire format must be reflected here, and adding a new
|
||||
// field requires updating MarshalPayload / unmarshalPayloadDetails together
|
||||
// with the field-uniqueness and wire-type checks in those functions.
|
||||
|
||||
syntax = "proto3";
|
||||
package nebula.handshake;
|
||||
|
||||
message NebulaHandshake {
|
||||
NebulaHandshakeDetails Details = 1;
|
||||
bytes Hmac = 2;
|
||||
}
|
||||
|
||||
message NebulaHandshakeDetails {
|
||||
bytes Cert = 1;
|
||||
uint32 InitiatorIndex = 2;
|
||||
uint32 ResponderIndex = 3;
|
||||
// Cookie was reserved for an anti-DoS mechanism that was never
|
||||
// implemented. No released version of nebula has ever populated it; the
|
||||
// hand-written parser silently skips it on read.
|
||||
uint64 Cookie = 4 [deprecated = true];
|
||||
uint64 Time = 5;
|
||||
uint32 CertVersion = 8;
|
||||
// reserved for WIP multiport
|
||||
reserved 6, 7;
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// testCertState holds cert material for a test peer.
|
||||
type testCertState struct {
|
||||
version cert.Version
|
||||
creds map[cert.Version]*Credential
|
||||
}
|
||||
|
||||
func (s *testCertState) getCredential(v cert.Version) *Credential {
|
||||
return s.creds[v]
|
||||
}
|
||||
|
||||
func newTestCertState(
|
||||
t *testing.T, ca cert.Certificate, caKey []byte, name string, networks []netip.Prefix,
|
||||
) *testCertState {
|
||||
return newTestCertStateWithCipher(t, ca, caKey, name, networks, noise.CipherChaChaPoly)
|
||||
}
|
||||
|
||||
func newTestCertStateWithCipher(
|
||||
t *testing.T, ca cert.Certificate, caKey []byte, name string, networks []netip.Prefix,
|
||||
cipher noise.CipherFunc,
|
||||
) *testCertState {
|
||||
t.Helper()
|
||||
c, _, rawPrivKey, _ := ct.NewTestCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
||||
name, ca.NotBefore(), ca.NotAfter(), networks, nil, nil,
|
||||
)
|
||||
|
||||
priv, _, _, err := cert.UnmarshalPrivateKeyFromPEM(rawPrivKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
hsBytes, err := c.MarshalForHandshakes()
|
||||
require.NoError(t, err)
|
||||
|
||||
ncs := noise.NewCipherSuite(noise.DH25519, cipher, noise.HashSHA256)
|
||||
return &testCertState{
|
||||
version: cert.Version2,
|
||||
creds: map[cert.Version]*Credential{
|
||||
cert.Version2: NewCredential(c, hsBytes, priv, ncs),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func testVerifier(pool *cert.CAPool) CertVerifier {
|
||||
return func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
||||
return pool.VerifyCertificate(time.Now(), c)
|
||||
}
|
||||
}
|
||||
|
||||
func newTestMachine(
|
||||
t *testing.T,
|
||||
cs *testCertState,
|
||||
verifier CertVerifier,
|
||||
initiator bool,
|
||||
localIndex uint32,
|
||||
) *Machine {
|
||||
t.Helper()
|
||||
m, err := NewMachine(
|
||||
cs.version, cs.getCredential,
|
||||
verifier, func() (uint32, error) { return localIndex, nil },
|
||||
initiator, header.HandshakeIXPSK0,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
return m
|
||||
}
|
||||
|
||||
func initiateHandshake(
|
||||
t *testing.T,
|
||||
initCS *testCertState, initVerifier CertVerifier,
|
||||
respCS *testCertState, respVerifier CertVerifier,
|
||||
) (initM, respM *Machine, respResult *Result, resp []byte, err error) {
|
||||
t.Helper()
|
||||
initM = newTestMachine(t, initCS, initVerifier, true, 100)
|
||||
msg1, merr := initM.Initiate(nil)
|
||||
require.NoError(t, merr)
|
||||
|
||||
respM = newTestMachine(t, respCS, respVerifier, false, 200)
|
||||
resp, respResult, err = respM.ProcessPacket(nil, msg1)
|
||||
return
|
||||
}
|
||||
|
||||
func doFullHandshake(
|
||||
t *testing.T, initCS, respCS *testCertState, caPool *cert.CAPool,
|
||||
) (initResult, respResult *Result) {
|
||||
t.Helper()
|
||||
v := testVerifier(caPool)
|
||||
|
||||
initM := newTestMachine(t, initCS, v, true, 1000)
|
||||
respM := newTestMachine(t, respCS, v, false, 2000)
|
||||
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, respResult, err := respM.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, respResult)
|
||||
require.NotEmpty(t, resp)
|
||||
|
||||
_, initResult, err = initM.ProcessPacket(nil, resp)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, initResult)
|
||||
|
||||
return initResult, respResult
|
||||
}
|
||||
@@ -0,0 +1,446 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/header"
|
||||
)
|
||||
|
||||
// IndexAllocator is called by the Machine to allocate a local index for the
|
||||
// handshake. It is called at most once, when the first outgoing message that
|
||||
// carries a payload is built.
|
||||
//
|
||||
// Implementations MUST NOT return 0. Zero is reserved as a sentinel meaning
|
||||
// "no index assigned" on the wire and in the payload-presence checks. If an
|
||||
// allocator ever returned 0, a legitimate handshake's payload could be
|
||||
// indistinguishable from an empty one and would be rejected.
|
||||
type IndexAllocator func() (uint32, error)
|
||||
|
||||
// CertVerifier is called by the Machine after reconstructing the peer's
|
||||
// certificate from the handshake. The verifier performs all validation
|
||||
// (CA trust, expiry, policy checks, allow lists).
|
||||
type CertVerifier func(cert.Certificate) (*cert.CachedCertificate, error)
|
||||
|
||||
// Result contains the results of a successful handshake.
|
||||
// Returned by ProcessPacket when the handshake is complete.
|
||||
type Result struct {
|
||||
EKey *noise.CipherState
|
||||
DKey *noise.CipherState
|
||||
Cipher noise.CipherFunc // identifies which post-handshake CipherState the data plane should wrap EKey/DKey in
|
||||
MyCert cert.Certificate
|
||||
RemoteCert *cert.CachedCertificate
|
||||
RemoteIndex uint32
|
||||
LocalIndex uint32
|
||||
HandshakeTime uint64
|
||||
MessageIndex uint64 // number of messages exchanged during the handshake
|
||||
Initiator bool
|
||||
}
|
||||
|
||||
// Machine drives a Noise handshake through N messages. It handles Noise
|
||||
// protocol operations, certificate reconstruction, and payload encoding.
|
||||
// Certificate validation is delegated to the caller via CertVerifier.
|
||||
//
|
||||
// A Machine is not safe for concurrent use. The caller must ensure that
|
||||
// Initiate and ProcessPacket are not called concurrently.
|
||||
//
|
||||
// Error contract: when ProcessPacket or Initiate returns an error, callers
|
||||
// must check Failed() to decide what to do next. If Failed() is false the
|
||||
// underlying noise state was not advanced (the packet was rejected before
|
||||
// ReadMessage took effect, or the rejection is non-fatal like a stale
|
||||
// retransmit) and the Machine can accept another packet. If Failed() is
|
||||
// true the Machine is unrecoverable and the caller must abandon it.
|
||||
type Machine struct {
|
||||
hs *noise.HandshakeState
|
||||
getCred GetCredentialFunc
|
||||
allocIndex IndexAllocator
|
||||
verifier CertVerifier
|
||||
result *Result
|
||||
msgs []msgFlags
|
||||
myVersion cert.Version
|
||||
subtype header.MessageSubType
|
||||
indexAllocated bool
|
||||
remoteCertSet bool
|
||||
payloadSet bool
|
||||
failed bool
|
||||
}
|
||||
|
||||
// NewMachine creates a handshake state machine. The subtype determines both
|
||||
// the noise pattern and the per-message content layout. The credential for
|
||||
// `version` is fetched via getCred and used to seed the noise.HandshakeState.
|
||||
// IndexAllocator is called lazily when the first outgoing payload is built.
|
||||
func NewMachine(
|
||||
version cert.Version,
|
||||
getCred GetCredentialFunc,
|
||||
verifier CertVerifier,
|
||||
allocIndex IndexAllocator,
|
||||
initiator bool,
|
||||
subtype header.MessageSubType,
|
||||
) (*Machine, error) {
|
||||
info, err := subtypeInfoFor(subtype)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cred := getCred(version)
|
||||
if cred == nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrNoCredential, version)
|
||||
}
|
||||
|
||||
hs, err := cred.buildHandshakeState(initiator, info.pattern)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build noise state: %w", err)
|
||||
}
|
||||
|
||||
return &Machine{
|
||||
hs: hs,
|
||||
subtype: subtype,
|
||||
msgs: info.msgs,
|
||||
getCred: getCred,
|
||||
allocIndex: allocIndex,
|
||||
verifier: verifier,
|
||||
myVersion: version,
|
||||
result: &Result{
|
||||
Initiator: initiator,
|
||||
Cipher: cred.cipherSuite,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Failed returns true if the Machine is in an unrecoverable state.
|
||||
func (m *Machine) Failed() bool {
|
||||
return m.failed
|
||||
}
|
||||
|
||||
// Subtype returns the handshake subtype this Machine was built for.
|
||||
func (m *Machine) Subtype() header.MessageSubType {
|
||||
return m.subtype
|
||||
}
|
||||
|
||||
// MessageIndex returns the noise handshake message index, which equals the
|
||||
// wire counter of the most recently sent or received message.
|
||||
func (m *Machine) MessageIndex() int {
|
||||
return m.hs.MessageIndex()
|
||||
}
|
||||
|
||||
// requireComplete checks that both a peer cert and payload have been received.
|
||||
// Marks the machine as failed if not.
|
||||
func (m *Machine) requireComplete() error {
|
||||
if !m.payloadSet || !m.remoteCertSet {
|
||||
m.failed = true
|
||||
return ErrIncompleteHandshake
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// myMsgFlags returns the flags for the current outgoing message.
|
||||
func (m *Machine) myMsgFlags() msgFlags {
|
||||
idx := m.hs.MessageIndex()
|
||||
if idx < len(m.msgs) {
|
||||
return m.msgs[idx]
|
||||
}
|
||||
return msgFlags{}
|
||||
}
|
||||
|
||||
// peerMsgFlags returns the flags for the message we just read.
|
||||
func (m *Machine) peerMsgFlags() msgFlags {
|
||||
idx := m.hs.MessageIndex() - 1
|
||||
if idx >= 0 && idx < len(m.msgs) {
|
||||
return m.msgs[idx]
|
||||
}
|
||||
return msgFlags{}
|
||||
}
|
||||
|
||||
// Initiate produces the first handshake message. Only valid for initiators,
|
||||
// and must be called exactly once before ProcessPacket.
|
||||
//
|
||||
// out is a destination buffer the message is appended to and returned. Pass
|
||||
// nil to allocate fresh, or pass a re-used buffer sliced to length 0 (e.g.
|
||||
// buf[:0]) with sufficient capacity to avoid allocation.
|
||||
//
|
||||
// An error return may not indicate a fatal condition, check Failed() to
|
||||
// determine if the Machine can still be used.
|
||||
func (m *Machine) Initiate(out []byte) ([]byte, error) {
|
||||
if m.failed {
|
||||
return nil, ErrMachineFailed
|
||||
}
|
||||
if !m.result.Initiator {
|
||||
m.failed = true
|
||||
return nil, ErrInitiateOnResponder
|
||||
}
|
||||
if m.hs.MessageIndex() != 0 {
|
||||
m.failed = true
|
||||
return nil, ErrInitiateAlreadyCalled
|
||||
}
|
||||
|
||||
// At MessageIndex=0 with RemoteIndex still zero, buildResponse produces
|
||||
// header counter 1 and remote index 0, which is what the initial message needs.
|
||||
out, _, _, err := m.buildResponse(out)
|
||||
if err != nil {
|
||||
m.failed = true
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// ProcessPacket handles an incoming handshake message. It advances the Noise
|
||||
// state, validates the peer certificate via the verifier, and optionally
|
||||
// produces a response.
|
||||
//
|
||||
// out is a destination buffer the response is appended to and returned. Pass
|
||||
// nil to allocate fresh, or pass a re-used buffer sliced to length 0 (e.g.
|
||||
// buf[:0]) with sufficient capacity to avoid allocation. The returned slice
|
||||
// is nil when no outgoing message is produced (handshake complete on this
|
||||
// side, or final message of a multi-message pattern).
|
||||
//
|
||||
// Returns a non-nil Result when the handshake is complete.
|
||||
// An error return may not indicate a fatal condition, check Failed() to
|
||||
// determine if the Machine can still be used.
|
||||
func (m *Machine) ProcessPacket(out, packet []byte) ([]byte, *Result, error) {
|
||||
if m.failed {
|
||||
return nil, nil, ErrMachineFailed
|
||||
}
|
||||
if len(packet) < header.Len {
|
||||
return nil, nil, ErrPacketTooShort
|
||||
}
|
||||
// Reject packets whose subtype doesn't match the one this Machine was
|
||||
// built for. A pending handshake that suddenly receives a different
|
||||
// subtype on its index is either a stray packet that matched by chance
|
||||
// or a peer protocol violation; drop it without failing the Machine so
|
||||
// the legitimate retransmit can still complete.
|
||||
if header.MessageSubType(packet[1]) != m.subtype {
|
||||
return nil, nil, ErrSubtypeMismatch
|
||||
}
|
||||
if m.result.Initiator && m.hs.MessageIndex() == 0 {
|
||||
m.failed = true
|
||||
return nil, nil, ErrInitiateNotCalled
|
||||
}
|
||||
|
||||
// The (eKey, dKey) ordering here is correct for IX, where the initiator
|
||||
// completes the handshake by reading the responder's stage-2 message.
|
||||
// noise returns (cs1, cs2) where cs1 is the initiator->responder cipher.
|
||||
// For 3-message patterns where a responder finishes by reading the final
|
||||
// message, this ordering would be wrong; revisit when XX/pqIX lands.
|
||||
msg, eKey, dKey, err := m.hs.ReadMessage(nil, packet[header.Len:])
|
||||
if err != nil {
|
||||
// Noise ReadMessage failed. The noise library checkpoints and rolls back
|
||||
// on failure, so the Machine is still alive. The caller can retry with
|
||||
// a different packet.
|
||||
return nil, nil, fmt.Errorf("noise ReadMessage: %w", err)
|
||||
}
|
||||
|
||||
// From here on, noise state has advanced. Any error is fatal.
|
||||
flags := m.peerMsgFlags()
|
||||
|
||||
if err := m.processPayload(msg, flags); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
// If ReadMessage derived keys, the handshake is complete. Noise should
|
||||
// always produce both keys together; asymmetry is a protocol invariant
|
||||
// violation.
|
||||
if eKey != nil || dKey != nil {
|
||||
if eKey == nil || dKey == nil {
|
||||
m.failed = true
|
||||
return nil, nil, ErrAsymmetricCipherKeys
|
||||
}
|
||||
if err := m.requireComplete(); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return nil, m.completed(eKey, dKey), nil
|
||||
}
|
||||
|
||||
// ReadMessage didn't complete, produce the next outgoing message
|
||||
out, dk, ek, err := m.buildResponse(out)
|
||||
if err != nil {
|
||||
m.failed = true
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
if ek != nil || dk != nil {
|
||||
if ek == nil || dk == nil {
|
||||
m.failed = true
|
||||
return nil, nil, ErrAsymmetricCipherKeys
|
||||
}
|
||||
if err := m.requireComplete(); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
return out, m.completed(ek, dk), nil
|
||||
}
|
||||
|
||||
return out, nil, nil
|
||||
}
|
||||
|
||||
func (m *Machine) completed(eKey, dKey *noise.CipherState) *Result {
|
||||
m.result.EKey = eKey
|
||||
m.result.DKey = dKey
|
||||
m.result.MessageIndex = uint64(m.hs.MessageIndex())
|
||||
return m.result
|
||||
}
|
||||
|
||||
func (m *Machine) processPayload(msg []byte, flags msgFlags) error {
|
||||
if len(msg) == 0 {
|
||||
if flags.expectsPayload || flags.expectsCert {
|
||||
m.failed = true
|
||||
return ErrMissingContent
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
payload, err := UnmarshalPayload(msg)
|
||||
if err != nil {
|
||||
m.failed = true
|
||||
return fmt.Errorf("unmarshal handshake: %w", err)
|
||||
}
|
||||
|
||||
// Assert the payload contains exactly what we expect
|
||||
hasPayloadData := payload.InitiatorIndex != 0 || payload.ResponderIndex != 0 || payload.Time != 0
|
||||
if hasPayloadData != flags.expectsPayload {
|
||||
m.failed = true
|
||||
return ErrUnexpectedContent
|
||||
}
|
||||
|
||||
hasCertData := len(payload.Cert) > 0
|
||||
if hasCertData != flags.expectsCert {
|
||||
m.failed = true
|
||||
return ErrUnexpectedContent
|
||||
}
|
||||
|
||||
// Process payload
|
||||
if flags.expectsPayload {
|
||||
if m.result.Initiator {
|
||||
m.result.RemoteIndex = payload.ResponderIndex
|
||||
} else {
|
||||
m.result.RemoteIndex = payload.InitiatorIndex
|
||||
}
|
||||
m.result.HandshakeTime = payload.Time
|
||||
m.payloadSet = true
|
||||
}
|
||||
|
||||
// Process certificate
|
||||
if flags.expectsCert {
|
||||
if err := m.validateCert(payload); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Machine) validateCert(payload Payload) error {
|
||||
cred := m.getCred(m.myVersion)
|
||||
if cred == nil {
|
||||
m.failed = true
|
||||
return fmt.Errorf("%w: %v", ErrNoCredential, m.myVersion)
|
||||
}
|
||||
rc, err := cert.Recombine(
|
||||
cert.Version(payload.CertVersion),
|
||||
payload.Cert,
|
||||
m.hs.PeerStatic(),
|
||||
cred.Cert.Curve(),
|
||||
)
|
||||
if err != nil {
|
||||
m.failed = true
|
||||
return fmt.Errorf("recombine cert: %w", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(rc.PublicKey(), m.hs.PeerStatic()) {
|
||||
m.failed = true
|
||||
return ErrPublicKeyMismatch
|
||||
}
|
||||
|
||||
// Version negotiation, if the peer sent a different version and we have it, switch
|
||||
if rc.Version() != m.myVersion {
|
||||
if m.getCred(rc.Version()) != nil {
|
||||
m.myVersion = rc.Version()
|
||||
}
|
||||
}
|
||||
|
||||
verified, err := m.verifier(rc)
|
||||
if err != nil {
|
||||
m.failed = true
|
||||
return fmt.Errorf("verify cert: %w", err)
|
||||
}
|
||||
|
||||
m.result.RemoteCert = verified
|
||||
m.remoteCertSet = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Machine) marshalOutgoing(flags msgFlags) ([]byte, error) {
|
||||
if !flags.expectsPayload && !flags.expectsCert {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
var p Payload
|
||||
if flags.expectsPayload {
|
||||
if !m.indexAllocated {
|
||||
index, err := m.allocIndex()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrIndexAllocation, err)
|
||||
}
|
||||
m.result.LocalIndex = index
|
||||
m.indexAllocated = true
|
||||
}
|
||||
|
||||
if m.result.Initiator {
|
||||
p.InitiatorIndex = m.result.LocalIndex
|
||||
} else {
|
||||
p.ResponderIndex = m.result.LocalIndex
|
||||
p.InitiatorIndex = m.result.RemoteIndex
|
||||
}
|
||||
p.Time = uint64(time.Now().UnixNano())
|
||||
}
|
||||
if flags.expectsCert {
|
||||
cred := m.getCred(m.myVersion)
|
||||
if cred == nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrNoCredential, m.myVersion)
|
||||
}
|
||||
p.Cert = cred.Bytes
|
||||
p.CertVersion = uint32(cred.Cert.Version())
|
||||
m.result.MyCert = cred.Cert
|
||||
}
|
||||
|
||||
return MarshalPayload(nil, p), nil
|
||||
}
|
||||
|
||||
func (m *Machine) buildResponse(out []byte) ([]byte, *noise.CipherState, *noise.CipherState, error) {
|
||||
flags := m.myMsgFlags()
|
||||
hsBytes, err := m.marshalOutgoing(flags)
|
||||
if err != nil {
|
||||
return nil, nil, nil, err
|
||||
}
|
||||
|
||||
// Extend out by header.Len to make room for the header. slices.Grow is a
|
||||
// no-op when the cap is already sufficient (the zero-copy case where the
|
||||
// caller passed a pre-sized buffer). header.Encode overwrites the new
|
||||
// bytes, so they don't need to be zeroed.
|
||||
start := len(out)
|
||||
out = slices.Grow(out, header.Len)[:start+header.Len]
|
||||
header.Encode(
|
||||
out[start:],
|
||||
header.Version, header.Handshake, m.subtype,
|
||||
m.result.RemoteIndex,
|
||||
uint64(m.hs.MessageIndex()+1),
|
||||
)
|
||||
|
||||
// noise.WriteMessage appends the encrypted handshake message to out,
|
||||
// reusing capacity when present.
|
||||
//
|
||||
// The (dKey, eKey) ordering here is correct for IX, where the responder
|
||||
// completes the handshake by writing the stage-2 message. noise returns
|
||||
// (cs1, cs2) where cs1 is the initiator->responder cipher (which is the
|
||||
// responder's decrypt key). For 3-message patterns where an initiator
|
||||
// finishes by writing the final message, this ordering would be wrong;
|
||||
// revisit when XX/pqIX lands.
|
||||
out, dKey, eKey, err := m.hs.WriteMessage(out, hsBytes)
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("noise WriteMessage: %w", err)
|
||||
}
|
||||
|
||||
return out, dKey, eKey, nil
|
||||
}
|
||||
@@ -0,0 +1,662 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
ct "github.com/slackhq/nebula/cert_test"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestMachineIXHappyPath(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
|
||||
initCS := newTestCertState(t, ca, caKey, "initiator", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
respCS := newTestCertState(t, ca, caKey, "responder", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
|
||||
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
||||
|
||||
assert.Equal(t, "responder", initR.RemoteCert.Certificate.Name())
|
||||
assert.Equal(t, "initiator", respR.RemoteCert.Certificate.Name())
|
||||
|
||||
assert.Equal(t, uint32(1000), initR.LocalIndex)
|
||||
assert.Equal(t, uint32(2000), initR.RemoteIndex)
|
||||
assert.Equal(t, uint32(2000), respR.LocalIndex)
|
||||
assert.Equal(t, uint32(1000), respR.RemoteIndex)
|
||||
|
||||
assert.Equal(t, uint64(2), initR.MessageIndex, "IX has 2 messages")
|
||||
assert.Equal(t, uint64(2), respR.MessageIndex, "IX has 2 messages")
|
||||
|
||||
ct1, err := initR.EKey.Encrypt(nil, nil, []byte("hello"))
|
||||
require.NoError(t, err)
|
||||
pt1, err := respR.DKey.Decrypt(nil, nil, ct1)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("hello"), pt1)
|
||||
|
||||
ct2, err := respR.EKey.Encrypt(nil, nil, []byte("world"))
|
||||
require.NoError(t, err)
|
||||
pt2, err := initR.DKey.Decrypt(nil, nil, ct2)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("world"), pt2)
|
||||
}
|
||||
|
||||
func TestMachineInitiateErrors(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
v := testVerifier(caPool)
|
||||
|
||||
t.Run("initiate on responder", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
_, err := m.Initiate(nil)
|
||||
require.ErrorIs(t, err, ErrInitiateOnResponder)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("initiate called twice", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, true, 100)
|
||||
_, err := m.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
_, err = m.Initiate(nil)
|
||||
require.ErrorIs(t, err, ErrInitiateAlreadyCalled)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("process packet before initiate on initiator", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, true, 100)
|
||||
_, _, err := m.ProcessPacket(nil, make([]byte, 100))
|
||||
require.ErrorIs(t, err, ErrInitiateNotCalled)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("calling failed machine", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
_, err := m.Initiate(nil) // fails: responder
|
||||
require.Error(t, err)
|
||||
_, err = m.Initiate(nil) // fails: already failed
|
||||
require.ErrorIs(t, err, ErrMachineFailed)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMachineProcessPacketErrors(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
v := testVerifier(caPool)
|
||||
|
||||
t.Run("packet too short", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
_, _, err := m.ProcessPacket(nil, []byte{1, 2, 3})
|
||||
require.ErrorIs(t, err, ErrPacketTooShort)
|
||||
assert.False(t, m.Failed(), "short packet should not kill machine")
|
||||
})
|
||||
|
||||
t.Run("noise decryption failure is recoverable", func(t *testing.T) {
|
||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
initM := newTestMachine(t, initCS, v, true, 100)
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
respM := newTestMachine(t, cs, v, false, 200)
|
||||
resp, _, err := respM.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
|
||||
corrupted := make([]byte, len(resp))
|
||||
copy(corrupted, resp)
|
||||
for i := header.Len; i < len(corrupted); i++ {
|
||||
corrupted[i] ^= 0xff
|
||||
}
|
||||
_, _, err = initM.ProcessPacket(nil, corrupted)
|
||||
require.Error(t, err)
|
||||
assert.False(t, initM.Failed(), "noise failure should be recoverable")
|
||||
|
||||
// And the machine should still complete a real handshake afterward.
|
||||
_, result, err := initM.ProcessPacket(nil, resp)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result, "initiator should complete on the legitimate response")
|
||||
})
|
||||
|
||||
t.Run("invalid cert is fatal", func(t *testing.T) {
|
||||
otherCA, _, otherCAKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
otherCS := newTestCertState(t, otherCA, otherCAKey, "other", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
|
||||
initM := newTestMachine(t, otherCS, testVerifier(ct.NewTestCAPool(otherCA)), true, 100)
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
respM := newTestMachine(t, cs, v, false, 200)
|
||||
_, _, err = respM.ProcessPacket(nil, msg1)
|
||||
require.Error(t, err)
|
||||
assert.True(t, respM.Failed(), "cert validation failure should kill machine")
|
||||
})
|
||||
|
||||
t.Run("subtype mismatch is recoverable", func(t *testing.T) {
|
||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
initM := newTestMachine(t, initCS, v, true, 100)
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Mutate the subtype byte (offset 1 in the header) to a value the
|
||||
// responder Machine wasn't built for.
|
||||
bad := make([]byte, len(msg1))
|
||||
copy(bad, msg1)
|
||||
bad[1] = 0xff
|
||||
|
||||
respM := newTestMachine(t, cs, v, false, 200)
|
||||
_, _, err = respM.ProcessPacket(nil, bad)
|
||||
require.ErrorIs(t, err, ErrSubtypeMismatch)
|
||||
assert.False(t, respM.Failed(), "subtype mismatch should not kill the machine")
|
||||
|
||||
// And the machine should still complete a real handshake afterward.
|
||||
resp, result, err := respM.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result, "responder should complete on the legitimate stage-1 packet")
|
||||
assert.NotEmpty(t, resp, "responder should produce a stage-2 reply")
|
||||
})
|
||||
}
|
||||
|
||||
// TestMachineProcessPayload exercises processPayload's internal validation
|
||||
// directly. Most of these failure modes can't be reached black-box once the
|
||||
// subtype check at the top of ProcessPacket gates external callers, so we
|
||||
// drive them by hand here for coverage.
|
||||
func TestMachineProcessPayload(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
v := testVerifier(caPool)
|
||||
|
||||
t.Run("empty message with expects fails", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
err := m.processPayload(nil, msgFlags{expectsPayload: true, expectsCert: true})
|
||||
require.ErrorIs(t, err, ErrMissingContent)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("empty message with no expects passes", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
err := m.processPayload(nil, msgFlags{})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("malformed protobuf is fatal", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
err := m.processPayload([]byte{0xff, 0xff, 0xff}, msgFlags{expectsPayload: true, expectsCert: true})
|
||||
require.Error(t, err)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("unexpected payload data is fatal", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
// A payload with index data when none was expected.
|
||||
bytes := MarshalPayload(nil, Payload{InitiatorIndex: 42, Time: 1})
|
||||
err := m.processPayload(bytes, msgFlags{expectsPayload: false, expectsCert: false})
|
||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("unexpected cert data is fatal", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
// A payload with cert when none was expected.
|
||||
bytes := MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2})
|
||||
err := m.processPayload(bytes, msgFlags{expectsPayload: false, expectsCert: false})
|
||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("missing payload data when expected is fatal", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
// Cert present, but no index/time fields.
|
||||
bytes := MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2})
|
||||
err := m.processPayload(bytes, msgFlags{expectsPayload: true, expectsCert: true})
|
||||
require.ErrorIs(t, err, ErrUnexpectedContent)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
}
|
||||
|
||||
// TestMachineRequireComplete checks the fail-on-incomplete-handshake path
|
||||
// directly. Like processPayload above this isn't reachable from a normal IX
|
||||
// flow, so we drive it by hand.
|
||||
func TestMachineRequireComplete(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
cs := newTestCertState(t, ca, caKey, "test", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
v := testVerifier(caPool)
|
||||
|
||||
t.Run("missing both fails", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
err := m.requireComplete()
|
||||
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("payload only fails", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
m.payloadSet = true
|
||||
err := m.requireComplete()
|
||||
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("cert only fails", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
m.remoteCertSet = true
|
||||
err := m.requireComplete()
|
||||
require.ErrorIs(t, err, ErrIncompleteHandshake)
|
||||
assert.True(t, m.Failed())
|
||||
})
|
||||
|
||||
t.Run("both set passes", func(t *testing.T) {
|
||||
m := newTestMachine(t, cs, v, false, 100)
|
||||
m.payloadSet = true
|
||||
m.remoteCertSet = true
|
||||
err := m.requireComplete()
|
||||
require.NoError(t, err)
|
||||
assert.False(t, m.Failed())
|
||||
})
|
||||
}
|
||||
|
||||
func TestMachineAESCipher(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
|
||||
initCS := newTestCertStateWithCipher(
|
||||
t, ca, caKey, "init",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||
noiseutil.CipherAESGCM,
|
||||
)
|
||||
respCS := newTestCertStateWithCipher(
|
||||
t, ca, caKey, "resp",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||
noiseutil.CipherAESGCM,
|
||||
)
|
||||
|
||||
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
||||
|
||||
ct1, err := initR.EKey.Encrypt(nil, nil, []byte("works"))
|
||||
require.NoError(t, err)
|
||||
pt1, err := respR.DKey.Decrypt(nil, nil, ct1)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("works"), pt1)
|
||||
|
||||
ct2, err := respR.EKey.Encrypt(nil, nil, []byte("back"))
|
||||
require.NoError(t, err)
|
||||
pt2, err := initR.DKey.Decrypt(nil, nil, ct2)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("back"), pt2)
|
||||
}
|
||||
|
||||
func TestResultFields(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
|
||||
initR, respR := doFullHandshake(t, initCS, respCS, caPool)
|
||||
|
||||
assert.True(t, initR.Initiator)
|
||||
assert.False(t, respR.Initiator)
|
||||
assert.NotZero(t, initR.HandshakeTime)
|
||||
assert.NotZero(t, respR.HandshakeTime)
|
||||
assert.NotNil(t, initR.RemoteCert)
|
||||
assert.NotNil(t, respR.RemoteCert)
|
||||
}
|
||||
|
||||
func TestMachineBufferReuse(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
v := testVerifier(caPool)
|
||||
|
||||
initM := newTestMachine(t, initCS, v, true, 1000)
|
||||
respM := newTestMachine(t, respCS, v, false, 2000)
|
||||
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("response writes into provided buffer", func(t *testing.T) {
|
||||
buf := make([]byte, 0, 4096)
|
||||
resp, result, err := respM.ProcessPacket(buf, msg1)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
|
||||
assert.NotEmpty(t, resp, "response should have content")
|
||||
assert.Equal(t, &buf[:1][0], &resp[:1][0],
|
||||
"response should reuse the provided buffer's backing array")
|
||||
})
|
||||
|
||||
t.Run("initiate writes into provided buffer", func(t *testing.T) {
|
||||
initM2 := newTestMachine(t, initCS, v, true, 3000)
|
||||
buf := make([]byte, 0, 4096)
|
||||
msg, err := initM2.Initiate(buf)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.NotEmpty(t, msg, "initiate should have content")
|
||||
assert.Equal(t, &buf[:1][0], &msg[:1][0],
|
||||
"initiate should reuse the provided buffer's backing array")
|
||||
})
|
||||
|
||||
t.Run("nil out still works", func(t *testing.T) {
|
||||
initM2 := newTestMachine(t, initCS, v, true, 4000)
|
||||
respM2 := newTestMachine(t, respCS, v, false, 5000)
|
||||
|
||||
msg1, err := initM2.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp, _, err := respM2.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
|
||||
out, result, err := initM2.ProcessPacket(nil, resp)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, result)
|
||||
assert.Nil(t, out, "initiator should have no response for IX msg2")
|
||||
})
|
||||
}
|
||||
|
||||
func TestMachineMsgIndexTracking(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
v := testVerifier(caPool)
|
||||
|
||||
initM := newTestMachine(t, initCS, v, true, 100)
|
||||
respM := newTestMachine(t, respCS, v, false, 200)
|
||||
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
resp1, result1, err := respM.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, result1)
|
||||
|
||||
_, result2, err := initM.ProcessPacket(nil, resp1)
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, result2)
|
||||
}
|
||||
|
||||
func TestMachineThreeMessagePattern(t *testing.T) {
|
||||
registerTestXXInfo(t)
|
||||
|
||||
// Use HandshakeXX (3 messages) to verify the Machine handles multi-message
|
||||
// patterns correctly. XX flow:
|
||||
// msg1 (I->R): [E] - payload only, no cert
|
||||
// msg2 (R->I): [E, ee, S, es] - payload + cert
|
||||
// msg3 (I->R): [S, se] - cert only (no payload, not first two)
|
||||
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
v := testVerifier(caPool)
|
||||
|
||||
initCS := newTestCertState(t, ca, caKey, "init", []netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")})
|
||||
respCS := newTestCertState(t, ca, caKey, "resp", []netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")})
|
||||
|
||||
initM, err := NewMachine(
|
||||
cert.Version2,
|
||||
initCS.getCredential, v,
|
||||
func() (uint32, error) { return 1000, nil },
|
||||
true, header.HandshakeXXPSK0,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
respM, err := NewMachine(
|
||||
cert.Version2,
|
||||
respCS.getCredential, v,
|
||||
func() (uint32, error) { return 2000, nil },
|
||||
false, header.HandshakeXXPSK0,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
// msg1: initiator -> responder (E only, no cert)
|
||||
msg1, err := initM.Initiate(nil)
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, msg1)
|
||||
|
||||
// Responder processes msg1, should not complete yet, should produce msg2
|
||||
msg2, result, err := respM.ProcessPacket(nil, msg1)
|
||||
require.NoError(t, err)
|
||||
assert.Nil(t, result, "XX should not complete on msg1")
|
||||
assert.NotEmpty(t, msg2, "responder should produce msg2")
|
||||
|
||||
// Initiator processes msg2: gets responder's cert, produces msg3, and
|
||||
// completes (WriteMessage for msg3 derives keys)
|
||||
msg3, initResult, err := initM.ProcessPacket(nil, msg2)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, initResult, "XX initiator should complete after reading msg2 and writing msg3")
|
||||
assert.NotEmpty(t, msg3, "initiator should produce msg3")
|
||||
assert.Equal(t, "resp", initResult.RemoteCert.Certificate.Name())
|
||||
|
||||
// Responder processes msg3: gets initiator's cert and completes
|
||||
_, respResult, err := respM.ProcessPacket(nil, msg3)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, respResult, "XX responder should complete on msg3")
|
||||
assert.Equal(t, "init", respResult.RemoteCert.Certificate.Name())
|
||||
|
||||
assert.Equal(t, uint64(3), initResult.MessageIndex, "XX has 3 messages")
|
||||
assert.Equal(t, uint64(3), respResult.MessageIndex, "XX has 3 messages")
|
||||
|
||||
// Verify keys work
|
||||
ct1, err := initResult.EKey.Encrypt(nil, nil, []byte("three messages"))
|
||||
require.NoError(t, err)
|
||||
pt1, err := respResult.DKey.Decrypt(nil, nil, ct1)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []byte("three messages"), pt1)
|
||||
}
|
||||
|
||||
// NOTE: ErrIncompleteHandshake is tested implicitly. It can't be triggered with
|
||||
// IX since the cert is always in the payload. A 3-message pattern test (HybridIX)
|
||||
// should exercise the case where cert arrives in msg3 and verify that completing
|
||||
// without it fails.
|
||||
|
||||
func TestMachineExpiredCert(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519,
|
||||
time.Now().Add(-24*time.Hour), time.Now().Add(24*time.Hour),
|
||||
nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
|
||||
expCert, _, expKeyPEM, _ := ct.NewTestCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, ca, caKey,
|
||||
"expired", time.Now().Add(-2*time.Hour), time.Now().Add(-1*time.Hour),
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")}, nil, nil,
|
||||
)
|
||||
expKey, _, _, err := cert.UnmarshalPrivateKeyFromPEM(expKeyPEM)
|
||||
require.NoError(t, err)
|
||||
expHsBytes, err := expCert.MarshalForHandshakes()
|
||||
require.NoError(t, err)
|
||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||
|
||||
expiredCS := &testCertState{
|
||||
version: cert.Version2,
|
||||
creds: map[cert.Version]*Credential{
|
||||
cert.Version2: NewCredential(expCert, expHsBytes, expKey, ncs),
|
||||
},
|
||||
}
|
||||
|
||||
respCS := newTestCertState(
|
||||
t, ca, caKey, "responder",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||
)
|
||||
|
||||
_, respM, _, _, err := initiateHandshake(
|
||||
t, expiredCS, testVerifier(caPool),
|
||||
respCS, testVerifier(caPool),
|
||||
)
|
||||
require.ErrorContains(t, err, "verify cert")
|
||||
assert.True(t, respM.Failed())
|
||||
}
|
||||
|
||||
func TestMachineNoCertNetworks(t *testing.T) {
|
||||
ca, _, caKey, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca)
|
||||
|
||||
caHsBytes, err := ca.MarshalForHandshakes()
|
||||
require.NoError(t, err)
|
||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||
|
||||
noNetCS := &testCertState{
|
||||
version: cert.Version2,
|
||||
creds: map[cert.Version]*Credential{
|
||||
cert.Version2: NewCredential(ca, caHsBytes, caKey, ncs),
|
||||
},
|
||||
}
|
||||
|
||||
respCS := newTestCertState(
|
||||
t, ca, caKey, "responder",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||
)
|
||||
|
||||
_, respM, _, _, err := initiateHandshake(
|
||||
t, noNetCS, testVerifier(caPool),
|
||||
respCS, testVerifier(caPool),
|
||||
)
|
||||
require.Error(t, err)
|
||||
assert.True(t, respM.Failed())
|
||||
}
|
||||
|
||||
func TestMachineDifferentCAs(t *testing.T) {
|
||||
ca1, _, caKey1, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
ca2, _, caKey2, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
|
||||
initCS := newTestCertState(
|
||||
t, ca1, caKey1, "init",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||
)
|
||||
respCS := newTestCertState(
|
||||
t, ca2, caKey2, "resp",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")},
|
||||
)
|
||||
|
||||
_, respM, _, _, err := initiateHandshake(
|
||||
t, initCS, testVerifier(ct.NewTestCAPool(ca1)),
|
||||
respCS, testVerifier(ct.NewTestCAPool(ca2)),
|
||||
)
|
||||
require.ErrorContains(t, err, "verify cert")
|
||||
assert.True(t, respM.Failed())
|
||||
}
|
||||
|
||||
func TestMachineVersionNegotiation(t *testing.T) {
|
||||
ca1, _, caKey1, _ := ct.NewTestCaCert(
|
||||
cert.Version1, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
ca2, _, caKey2, _ := ct.NewTestCaCert(
|
||||
cert.Version2, cert.Curve_CURVE25519, time.Time{}, time.Time{}, nil, nil, nil,
|
||||
)
|
||||
caPool := ct.NewTestCAPool(ca1, ca2)
|
||||
|
||||
makeMultiVersionResp := func(t *testing.T) *testCertState {
|
||||
t.Helper()
|
||||
respCertV1, _, respKeyPEM, _ := ct.NewTestCert(
|
||||
cert.Version1, cert.Curve_CURVE25519, ca1, caKey1, "resp",
|
||||
ca1.NotBefore(), ca1.NotAfter(),
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")}, nil, nil,
|
||||
)
|
||||
respKey, _, _, _ := cert.UnmarshalPrivateKeyFromPEM(respKeyPEM)
|
||||
respCertV2, _ := ct.NewTestCertDifferentVersion(respCertV1, cert.Version2, ca2, caKey2)
|
||||
respHsV1, _ := respCertV1.MarshalForHandshakes()
|
||||
respHsV2, _ := respCertV2.MarshalForHandshakes()
|
||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||
return &testCertState{
|
||||
version: cert.Version1,
|
||||
creds: map[cert.Version]*Credential{
|
||||
cert.Version1: NewCredential(respCertV1, respHsV1, respKey, ncs),
|
||||
cert.Version2: NewCredential(respCertV2, respHsV2, respKey, ncs),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("responder matches initiator version", func(t *testing.T) {
|
||||
initCS := newTestCertState(
|
||||
t, ca2, caKey2, "init",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||
)
|
||||
respCS := makeMultiVersionResp(t)
|
||||
v := testVerifier(caPool)
|
||||
|
||||
initM, _, respResult, resp, err := initiateHandshake(
|
||||
t, initCS, v,
|
||||
respCS, v,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, respResult)
|
||||
|
||||
assert.Equal(t, cert.Version2, respResult.MyCert.Version(),
|
||||
"responder should negotiate to initiator's version")
|
||||
|
||||
_, initResult, err := initM.ProcessPacket(nil, resp)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, initResult)
|
||||
assert.Equal(t, cert.Version2, initResult.RemoteCert.Certificate.Version(),
|
||||
"initiator should see V2 cert from responder")
|
||||
})
|
||||
|
||||
t.Run("responder keeps version when no match available", func(t *testing.T) {
|
||||
initCS := newTestCertState(
|
||||
t, ca2, caKey2, "init",
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.1/24")},
|
||||
)
|
||||
|
||||
respCert, _, respKeyPEM, _ := ct.NewTestCert(
|
||||
cert.Version1, cert.Curve_CURVE25519, ca1, caKey1, "resp",
|
||||
ca1.NotBefore(), ca1.NotAfter(),
|
||||
[]netip.Prefix{netip.MustParsePrefix("10.0.0.2/24")}, nil, nil,
|
||||
)
|
||||
respKey, _, _, _ := cert.UnmarshalPrivateKeyFromPEM(respKeyPEM)
|
||||
respHs, _ := respCert.MarshalForHandshakes()
|
||||
ncs := noise.NewCipherSuite(noise.DH25519, noise.CipherChaChaPoly, noise.HashSHA256)
|
||||
respCS := &testCertState{
|
||||
version: cert.Version1,
|
||||
creds: map[cert.Version]*Credential{
|
||||
cert.Version1: NewCredential(respCert, respHs, respKey, ncs),
|
||||
},
|
||||
}
|
||||
|
||||
v := testVerifier(caPool)
|
||||
_, _, respResult, _, err := initiateHandshake(
|
||||
t, initCS, v,
|
||||
respCS, v,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, respResult)
|
||||
|
||||
assert.Equal(t, cert.Version1, respResult.MyCert.Version(),
|
||||
"responder should keep V1 when V2 not available")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/header"
|
||||
)
|
||||
|
||||
// msgFlags tracks what application data a handshake message carries.
|
||||
type msgFlags struct {
|
||||
expectsPayload bool // message carries indexes and time
|
||||
expectsCert bool // message carries the certificate
|
||||
}
|
||||
|
||||
// subtypeInfo bundles the noise pattern with the per-message flags for a
|
||||
// given handshake subtype.
|
||||
type subtypeInfo struct {
|
||||
pattern noise.HandshakePattern
|
||||
msgs []msgFlags
|
||||
}
|
||||
|
||||
// subtypeInfos defines the noise pattern and message content layout for each
|
||||
// handshake subtype.
|
||||
var subtypeInfos = map[header.MessageSubType]subtypeInfo{
|
||||
// IX: 2 messages, both carry payload and cert
|
||||
header.HandshakeIXPSK0: {
|
||||
pattern: noise.HandshakeIX,
|
||||
msgs: []msgFlags{
|
||||
{expectsPayload: true, expectsCert: true},
|
||||
{expectsPayload: true, expectsCert: true},
|
||||
},
|
||||
},
|
||||
|
||||
// XX: 3 messages
|
||||
// msg1 (I->R): payload only
|
||||
// msg2 (R->I): payload + cert
|
||||
// msg3 (I->R): cert only
|
||||
//header.HandshakeXXPSK0: {
|
||||
// pattern: noise.HandshakeXX,
|
||||
// msgs: []msgFlags{
|
||||
// {expectsPayload: true, expectsCert: false},
|
||||
// {expectsPayload: true, expectsCert: true},
|
||||
// {expectsPayload: false, expectsCert: true},
|
||||
// },
|
||||
//},
|
||||
}
|
||||
|
||||
func subtypeInfoFor(subtype header.MessageSubType) (subtypeInfo, error) {
|
||||
if info, ok := subtypeInfos[subtype]; ok {
|
||||
return info, nil
|
||||
}
|
||||
return subtypeInfo{}, fmt.Errorf("%w: %d", ErrUnknownSubtype, subtype)
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestSubtypeInfo(t *testing.T) {
|
||||
t.Run("IX", func(t *testing.T) {
|
||||
info, err := subtypeInfoFor(header.HandshakeIXPSK0)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, noise.HandshakeIX.Name, info.pattern.Name)
|
||||
require.Len(t, info.msgs, 2)
|
||||
// msg1: payload + cert
|
||||
assert.True(t, info.msgs[0].expectsPayload)
|
||||
assert.True(t, info.msgs[0].expectsCert)
|
||||
// msg2: payload + cert
|
||||
assert.True(t, info.msgs[1].expectsPayload)
|
||||
assert.True(t, info.msgs[1].expectsCert)
|
||||
})
|
||||
|
||||
t.Run("XX", func(t *testing.T) {
|
||||
registerTestXXInfo(t)
|
||||
info, err := subtypeInfoFor(header.HandshakeXXPSK0)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, noise.HandshakeXX.Name, info.pattern.Name)
|
||||
require.Len(t, info.msgs, 3)
|
||||
// msg1: payload only
|
||||
assert.True(t, info.msgs[0].expectsPayload)
|
||||
assert.False(t, info.msgs[0].expectsCert)
|
||||
// msg2: payload + cert
|
||||
assert.True(t, info.msgs[1].expectsPayload)
|
||||
assert.True(t, info.msgs[1].expectsCert)
|
||||
// msg3: cert only
|
||||
assert.False(t, info.msgs[2].expectsPayload)
|
||||
assert.True(t, info.msgs[2].expectsCert)
|
||||
})
|
||||
|
||||
t.Run("unknown subtype returns error", func(t *testing.T) {
|
||||
_, err := subtypeInfoFor(99)
|
||||
require.ErrorIs(t, err, ErrUnknownSubtype)
|
||||
})
|
||||
}
|
||||
|
||||
// registerTestXXInfo temporarily registers XX subtype info for testing.
|
||||
func registerTestXXInfo(t *testing.T) {
|
||||
t.Helper()
|
||||
subtypeInfos[header.HandshakeXXPSK0] = subtypeInfo{
|
||||
pattern: noise.HandshakeXX,
|
||||
msgs: []msgFlags{
|
||||
{expectsPayload: true, expectsCert: false},
|
||||
{expectsPayload: true, expectsCert: true},
|
||||
{expectsPayload: false, expectsCert: true},
|
||||
},
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
delete(subtypeInfos, header.HandshakeXXPSK0)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
)
|
||||
|
||||
var (
|
||||
errInvalidHandshakeMessage = errors.New("invalid handshake message")
|
||||
errInvalidHandshakeDetails = errors.New("invalid handshake details")
|
||||
)
|
||||
|
||||
// Payload represents the decoded fields of a handshake message.
|
||||
// Wire format is protobuf-compatible with NebulaHandshake{Details: NebulaHandshakeDetails{...}}.
|
||||
type Payload struct {
|
||||
Cert []byte
|
||||
InitiatorIndex uint32
|
||||
ResponderIndex uint32
|
||||
Time uint64
|
||||
CertVersion uint32
|
||||
}
|
||||
|
||||
// Proto field numbers for NebulaHandshakeDetails
|
||||
const (
|
||||
fieldCert = 1 // bytes
|
||||
fieldInitiatorIndex = 2 // uint32
|
||||
fieldResponderIndex = 3 // uint32
|
||||
fieldTime = 5 // uint64
|
||||
fieldCertVersion = 8 // uint32
|
||||
)
|
||||
|
||||
// MarshalPayload encodes a handshake payload in protobuf wire format compatible
|
||||
// with NebulaHandshake{Details: NebulaHandshakeDetails{...}}.
|
||||
// Returns out (which may be nil), with the marshalled Payload appended to it.
|
||||
func MarshalPayload(out []byte, p Payload) []byte {
|
||||
var details []byte
|
||||
|
||||
if len(p.Cert) > 0 {
|
||||
details = protowire.AppendTag(details, fieldCert, protowire.BytesType)
|
||||
details = protowire.AppendBytes(details, p.Cert)
|
||||
}
|
||||
if p.InitiatorIndex != 0 {
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, uint64(p.InitiatorIndex))
|
||||
}
|
||||
if p.ResponderIndex != 0 {
|
||||
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, uint64(p.ResponderIndex))
|
||||
}
|
||||
if p.Time != 0 {
|
||||
details = protowire.AppendTag(details, fieldTime, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, p.Time)
|
||||
}
|
||||
if p.CertVersion != 0 {
|
||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, uint64(p.CertVersion))
|
||||
}
|
||||
|
||||
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
||||
out = protowire.AppendBytes(out, details)
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
// UnmarshalPayload decodes a protobuf-encoded NebulaHandshake message.
|
||||
func UnmarshalPayload(b []byte) (Payload, error) {
|
||||
var p Payload
|
||||
|
||||
for len(b) > 0 {
|
||||
num, typ, n := protowire.ConsumeTag(b)
|
||||
if n < 0 {
|
||||
return p, errInvalidHandshakeMessage
|
||||
}
|
||||
b = b[n:]
|
||||
|
||||
switch {
|
||||
case num == 1 && typ == protowire.BytesType:
|
||||
details, n := protowire.ConsumeBytes(b)
|
||||
if n < 0 {
|
||||
return p, errInvalidHandshakeMessage
|
||||
}
|
||||
b = b[n:]
|
||||
if err := unmarshalPayloadDetails(&p, details); err != nil {
|
||||
return p, err
|
||||
}
|
||||
default:
|
||||
n := protowire.ConsumeFieldValue(num, typ, b)
|
||||
if n < 0 {
|
||||
return p, errInvalidHandshakeMessage
|
||||
}
|
||||
b = b[n:]
|
||||
}
|
||||
}
|
||||
|
||||
return p, nil
|
||||
}
|
||||
|
||||
func unmarshalPayloadDetails(p *Payload, b []byte) error {
|
||||
for len(b) > 0 {
|
||||
num, typ, n := protowire.ConsumeTag(b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
b = b[n:]
|
||||
|
||||
// For known field numbers, reject any non-matching wire type as a
|
||||
// hard error rather than silently skipping. The caller will catch
|
||||
// missing-field cases downstream, but a wire-type mismatch on a tag
|
||||
// we know is a peer protocol violation worth flagging here.
|
||||
// Repeated occurrences of a singular field follow proto3 last-wins.
|
||||
switch num {
|
||||
case fieldCert:
|
||||
if typ != protowire.BytesType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeBytes(b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.Cert = append([]byte(nil), v...)
|
||||
b = b[n:]
|
||||
case fieldInitiatorIndex:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.InitiatorIndex = uint32(v)
|
||||
b = b[n:]
|
||||
case fieldResponderIndex:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.ResponderIndex = uint32(v)
|
||||
b = b[n:]
|
||||
case fieldTime:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.Time = v
|
||||
b = b[n:]
|
||||
case fieldCertVersion:
|
||||
if typ != protowire.VarintType {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
v, n := protowire.ConsumeVarint(b)
|
||||
if n < 0 || v > math.MaxUint32 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
p.CertVersion = uint32(v)
|
||||
b = b[n:]
|
||||
default:
|
||||
n := protowire.ConsumeFieldValue(num, typ, b)
|
||||
if n < 0 {
|
||||
return errInvalidHandshakeDetails
|
||||
}
|
||||
b = b[n:]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,361 @@
|
||||
package handshake
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"math"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
)
|
||||
|
||||
func TestPayloadRoundTrip(t *testing.T) {
|
||||
t.Run("all fields set", func(t *testing.T) {
|
||||
data := MarshalPayload(nil, Payload{
|
||||
Cert: []byte("test-cert-bytes"),
|
||||
CertVersion: 2,
|
||||
InitiatorIndex: 12345,
|
||||
ResponderIndex: 67890,
|
||||
Time: 1234567890,
|
||||
})
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, []byte("test-cert-bytes"), got.Cert)
|
||||
assert.Equal(t, uint32(12345), got.InitiatorIndex)
|
||||
assert.Equal(t, uint32(67890), got.ResponderIndex)
|
||||
assert.Equal(t, uint64(1234567890), got.Time)
|
||||
assert.Equal(t, uint32(2), got.CertVersion)
|
||||
})
|
||||
|
||||
t.Run("minimal fields", func(t *testing.T) {
|
||||
data := MarshalPayload(nil, Payload{InitiatorIndex: 1})
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, uint32(1), got.InitiatorIndex)
|
||||
assert.Equal(t, uint32(0), got.ResponderIndex)
|
||||
assert.Equal(t, uint64(0), got.Time)
|
||||
assert.Nil(t, got.Cert)
|
||||
})
|
||||
|
||||
t.Run("empty payload", func(t *testing.T) {
|
||||
data := MarshalPayload(nil, Payload{})
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, uint32(0), got.InitiatorIndex)
|
||||
})
|
||||
|
||||
t.Run("large cert bytes", func(t *testing.T) {
|
||||
bigCert := make([]byte, 4096)
|
||||
for i := range bigCert {
|
||||
bigCert[i] = byte(i % 256)
|
||||
}
|
||||
|
||||
data := MarshalPayload(nil, Payload{
|
||||
Cert: bigCert,
|
||||
CertVersion: 2,
|
||||
InitiatorIndex: 999,
|
||||
})
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, bigCert, got.Cert)
|
||||
assert.Equal(t, uint32(999), got.InitiatorIndex)
|
||||
})
|
||||
|
||||
t.Run("append to existing buffer", func(t *testing.T) {
|
||||
prefix := []byte("prefix")
|
||||
data := MarshalPayload(prefix, Payload{InitiatorIndex: 42})
|
||||
|
||||
assert.Equal(t, []byte("prefix"), data[:6])
|
||||
|
||||
got, err := UnmarshalPayload(data[6:])
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPayloadUnknownFields(t *testing.T) {
|
||||
t.Run("unknown field in outer message is skipped", func(t *testing.T) {
|
||||
// Marshal a normal payload then append an unknown field (field 99, varint)
|
||||
data := MarshalPayload(nil, Payload{InitiatorIndex: 42})
|
||||
data = protowire.AppendTag(data, 99, protowire.VarintType)
|
||||
data = protowire.AppendVarint(data, 12345)
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
||||
})
|
||||
|
||||
t.Run("unknown field in details is skipped", func(t *testing.T) {
|
||||
// Build details with a known field + unknown field
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 77)
|
||||
// Unknown field 50, varint
|
||||
details = protowire.AppendTag(details, 50, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 9999)
|
||||
// Another known field after the unknown one
|
||||
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 88)
|
||||
|
||||
// Wrap in outer message
|
||||
var data []byte
|
||||
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
||||
data = protowire.AppendBytes(data, details)
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(77), got.InitiatorIndex)
|
||||
assert.Equal(t, uint32(88), got.ResponderIndex)
|
||||
})
|
||||
|
||||
t.Run("reserved fields 6 and 7 are skipped", func(t *testing.T) {
|
||||
// Fields 6 and 7 are reserved in the proto definition
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 100)
|
||||
details = protowire.AppendTag(details, 6, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 1)
|
||||
details = protowire.AppendTag(details, 7, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 2)
|
||||
|
||||
var data []byte
|
||||
data = protowire.AppendTag(data, 1, protowire.BytesType)
|
||||
data = protowire.AppendBytes(data, details)
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(100), got.InitiatorIndex)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPayloadBytesConsumed(t *testing.T) {
|
||||
t.Run("all bytes consumed on valid input", func(t *testing.T) {
|
||||
original := Payload{
|
||||
Cert: []byte("cert"),
|
||||
CertVersion: 2,
|
||||
InitiatorIndex: 100,
|
||||
ResponderIndex: 200,
|
||||
Time: 999,
|
||||
}
|
||||
data := MarshalPayload(nil, original)
|
||||
|
||||
got, err := UnmarshalPayload(data)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Re-marshal and compare — proves we consumed and reproduced all fields
|
||||
remarshaled := MarshalPayload(nil, got)
|
||||
assert.Equal(t, data, remarshaled)
|
||||
})
|
||||
}
|
||||
|
||||
// wrapDetails wraps raw detail bytes in the outer NebulaHandshake envelope
|
||||
// so UnmarshalPayload can reach unmarshalPayloadDetails.
|
||||
func wrapDetails(details []byte) []byte {
|
||||
var out []byte
|
||||
out = protowire.AppendTag(out, 1, protowire.BytesType)
|
||||
out = protowire.AppendBytes(out, details)
|
||||
return out
|
||||
}
|
||||
|
||||
func TestPayloadUnmarshalErrors(t *testing.T) {
|
||||
t.Run("nil input", func(t *testing.T) {
|
||||
got, err := UnmarshalPayload(nil)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(0), got.InitiatorIndex)
|
||||
})
|
||||
|
||||
t.Run("truncated outer tag", func(t *testing.T) {
|
||||
_, err := UnmarshalPayload([]byte{0x80})
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated outer details field", func(t *testing.T) {
|
||||
_, err := UnmarshalPayload([]byte{0x0a, 0x64, 0x01, 0x02, 0x03, 0x04, 0x05})
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated outer unknown field", func(t *testing.T) {
|
||||
// Valid tag for unknown field 99 varint, but no value follows
|
||||
var data []byte
|
||||
data = protowire.AppendTag(data, 99, protowire.VarintType)
|
||||
_, err := UnmarshalPayload(data)
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated details tag", func(t *testing.T) {
|
||||
_, err := UnmarshalPayload(wrapDetails([]byte{0x80}))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated cert bytes", func(t *testing.T) {
|
||||
// Field 1 (cert), bytes type, length 10 but only 2 bytes
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldCert, protowire.BytesType)
|
||||
details = append(details, 0x0a, 0x01, 0x02) // length 10, only 2 bytes
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated initiator index varint", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = append(details, 0x80) // incomplete varint
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated responder index varint", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldResponderIndex, protowire.VarintType)
|
||||
details = append(details, 0x80)
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated time varint", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldTime, protowire.VarintType)
|
||||
details = append(details, 0x80)
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated cert version varint", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||
details = append(details, 0x80)
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("truncated unknown field in details", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, 50, protowire.VarintType)
|
||||
details = append(details, 0x80) // incomplete varint
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("cert with wrong wire type rejected", func(t *testing.T) {
|
||||
// fieldCert as Varint instead of Bytes.
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldCert, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 42)
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("initiator index with wrong wire type rejected", func(t *testing.T) {
|
||||
// fieldInitiatorIndex as Bytes instead of Varint.
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.BytesType)
|
||||
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("time with wrong wire type rejected", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldTime, protowire.BytesType)
|
||||
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("cert version with wrong wire type rejected", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.BytesType)
|
||||
details = protowire.AppendBytes(details, []byte{1, 2, 3})
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("repeated singular field follows proto3 last-wins", func(t *testing.T) {
|
||||
// Per proto3, multiple instances of a singular field are accepted and
|
||||
// the last value wins. We keep this behavior so that peers using
|
||||
// alternative encoders aren't rejected.
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 1)
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, 42)
|
||||
got, err := UnmarshalPayload(wrapDetails(details))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, uint32(42), got.InitiatorIndex)
|
||||
})
|
||||
|
||||
t.Run("initiator index varint overflow rejected", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldInitiatorIndex, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, math.MaxUint32+1)
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("cert version varint overflow rejected", func(t *testing.T) {
|
||||
var details []byte
|
||||
details = protowire.AppendTag(details, fieldCertVersion, protowire.VarintType)
|
||||
details = protowire.AppendVarint(details, math.MaxUint32+1)
|
||||
_, err := UnmarshalPayload(wrapDetails(details))
|
||||
assert.Error(t, err)
|
||||
})
|
||||
|
||||
}
|
||||
|
||||
// FuzzPayload feeds arbitrary bytes through UnmarshalPayload to confirm it
|
||||
// never panics, and for any input that parses cleanly, that re-marshal +
|
||||
// re-parse is a fix-point. Inputs come from an authenticated peer (post-
|
||||
// noise-decrypt), so the threat model is "valid peer behaving arbitrarily,"
|
||||
// not "unauthenticated injection."
|
||||
func FuzzPayload(f *testing.F) {
|
||||
// Seed corpus with a handful of known-good shapes.
|
||||
f.Add(MarshalPayload(nil, Payload{}))
|
||||
f.Add(MarshalPayload(nil, Payload{Cert: []byte{1, 2, 3}, CertVersion: 2}))
|
||||
f.Add(MarshalPayload(nil, Payload{InitiatorIndex: 42, Time: 1}))
|
||||
f.Add(MarshalPayload(nil, Payload{
|
||||
Cert: []byte("seed-cert"),
|
||||
InitiatorIndex: 1,
|
||||
ResponderIndex: 2,
|
||||
Time: 3,
|
||||
CertVersion: 2,
|
||||
}))
|
||||
f.Add([]byte{})
|
||||
f.Add([]byte{0xff})
|
||||
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
p1, err := UnmarshalPayload(data)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
// For any input that parses, re-marshaling and re-parsing must
|
||||
// yield an equivalent Payload. This catches dispatch bugs (e.g.
|
||||
// emitting a field on marshal that we don't accept on parse) and
|
||||
// any non-idempotent parsing behavior.
|
||||
b2 := MarshalPayload(nil, p1)
|
||||
p2, err := UnmarshalPayload(b2)
|
||||
if err != nil {
|
||||
t.Fatalf("re-parse of self-marshaled payload failed: %v\nintermediate: %x\n", err, b2)
|
||||
}
|
||||
if !payloadsEqual(p1, p2) {
|
||||
t.Fatalf("re-marshal not idempotent\nfirst: %+v\nsecond: %+v", p1, p2)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func payloadsEqual(a, b Payload) bool {
|
||||
return bytes.Equal(a.Cert, b.Cert) &&
|
||||
a.InitiatorIndex == b.InitiatorIndex &&
|
||||
a.ResponderIndex == b.ResponderIndex &&
|
||||
a.Time == b.Time &&
|
||||
a.CertVersion == b.CertVersion
|
||||
}
|
||||
-678
@@ -1,678 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/netip"
|
||||
"time"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/header"
|
||||
)
|
||||
|
||||
// NOISE IX Handshakes
|
||||
|
||||
// This function constructs a handshake packet, but does not actually send it
|
||||
// Sending is done by the handshake manager
|
||||
func ixHandshakeStage0(f *Interface, hh *HandshakeHostInfo) bool {
|
||||
err := f.handshakeManager.allocateIndex(hh)
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).Error("Failed to generate index")
|
||||
return false
|
||||
}
|
||||
|
||||
cs := f.pki.getCertState()
|
||||
v := cs.initiatingVersion
|
||||
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
||||
v = hh.initiatingVersionOverride
|
||||
} else if v < cert.Version2 {
|
||||
// If we're connecting to a v6 address we should encourage use of a V2 cert
|
||||
for _, a := range hh.hostinfo.vpnAddrs {
|
||||
if a.Is6() {
|
||||
v = cert.Version2
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
crt := cs.getCertificate(v)
|
||||
if crt == nil {
|
||||
f.l.WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||
WithField("certVersion", v).
|
||||
Error("Unable to handshake with host because no certificate is available")
|
||||
return false
|
||||
}
|
||||
|
||||
crtHs := cs.getHandshakeBytes(v)
|
||||
if crtHs == nil {
|
||||
f.l.WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||
WithField("certVersion", v).
|
||||
Error("Unable to handshake with host because no certificate handshake bytes is available")
|
||||
return false
|
||||
}
|
||||
|
||||
ci, err := NewConnectionState(f.l, cs, crt, true, noise.HandshakeIX)
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||
WithField("certVersion", v).
|
||||
Error("Failed to create connection state")
|
||||
return false
|
||||
}
|
||||
hh.hostinfo.ConnectionState = ci
|
||||
|
||||
hs := &NebulaHandshake{
|
||||
Details: &NebulaHandshakeDetails{
|
||||
InitiatorIndex: hh.hostinfo.localIndexId,
|
||||
Time: uint64(time.Now().UnixNano()),
|
||||
Cert: crtHs,
|
||||
CertVersion: uint32(v),
|
||||
},
|
||||
}
|
||||
|
||||
hsBytes, err := hs.Marshal()
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||
WithField("certVersion", v).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).Error("Failed to marshal handshake message")
|
||||
return false
|
||||
}
|
||||
|
||||
h := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, 0, 1)
|
||||
|
||||
msg, _, _, err := ci.H.WriteMessage(h, hsBytes)
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hh.hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).Error("Failed to call noise.WriteMessage")
|
||||
return false
|
||||
}
|
||||
|
||||
// We are sending handshake packet 1, so we don't expect to receive
|
||||
// handshake packet 1 from the responder
|
||||
ci.window.Update(f.l, 1)
|
||||
|
||||
hh.hostinfo.HandshakePacket[0] = msg
|
||||
hh.ready = true
|
||||
return true
|
||||
}
|
||||
|
||||
func ixHandshakeStage1(f *Interface, via ViaSender, packet []byte, h *header.H) {
|
||||
cs := f.pki.getCertState()
|
||||
crt := cs.GetDefaultCertificate()
|
||||
if crt == nil {
|
||||
f.l.WithField("from", via).
|
||||
WithField("handshake", m{"stage": 0, "style": "ix_psk0"}).
|
||||
WithField("certVersion", cs.initiatingVersion).
|
||||
Error("Unable to handshake with host because no certificate is available")
|
||||
return
|
||||
}
|
||||
|
||||
ci, err := NewConnectionState(f.l, cs, crt, false, noise.HandshakeIX)
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Error("Failed to create connection state")
|
||||
return
|
||||
}
|
||||
|
||||
// Mark packet 1 as seen so it doesn't show up as missed
|
||||
ci.window.Update(f.l, 1)
|
||||
|
||||
msg, _, _, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Error("Failed to call noise.ReadMessage")
|
||||
return
|
||||
}
|
||||
|
||||
hs := &NebulaHandshake{}
|
||||
err = hs.Unmarshal(msg)
|
||||
if err != nil || hs.Details == nil {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Error("Failed unmarshal handshake message")
|
||||
return
|
||||
}
|
||||
|
||||
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Info("Handshake did not contain a certificate")
|
||||
return
|
||||
}
|
||||
|
||||
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
||||
if err != nil {
|
||||
fp, fperr := rc.Fingerprint()
|
||||
if fperr != nil {
|
||||
fp = "<error generating certificate fingerprint>"
|
||||
}
|
||||
|
||||
e := f.l.WithError(err).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
WithField("certVpnNetworks", rc.Networks()).
|
||||
WithField("certFingerprint", fp)
|
||||
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
e = e.WithField("cert", rc)
|
||||
}
|
||||
|
||||
e.Info("Invalid certificate from host")
|
||||
return
|
||||
}
|
||||
|
||||
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||
f.l.WithField("from", via).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
WithField("cert", remoteCert).Info("public key mismatch between certificate and handshake")
|
||||
return
|
||||
}
|
||||
|
||||
if remoteCert.Certificate.Version() != ci.myCert.Version() {
|
||||
// We started off using the wrong certificate version, lets see if we can match the version that was sent to us
|
||||
myCertOtherVersion := cs.getCertificate(remoteCert.Certificate.Version())
|
||||
if myCertOtherVersion == nil {
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithError(err).WithFields(m{
|
||||
"from": via,
|
||||
"handshake": m{"stage": 1, "style": "ix_psk0"},
|
||||
"cert": remoteCert,
|
||||
}).Debug("Might be unable to handshake with host due to missing certificate version")
|
||||
}
|
||||
} else {
|
||||
// Record the certificate we are actually using
|
||||
ci.myCert = myCertOtherVersion
|
||||
}
|
||||
}
|
||||
|
||||
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("cert", remoteCert).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Info("No networks in certificate")
|
||||
return
|
||||
}
|
||||
|
||||
certName := remoteCert.Certificate.Name()
|
||||
certVersion := remoteCert.Certificate.Version()
|
||||
fingerprint := remoteCert.Fingerprint
|
||||
issuer := remoteCert.Certificate.Issuer()
|
||||
vpnNetworks := remoteCert.Certificate.Networks()
|
||||
|
||||
anyVpnAddrsInCommon := false
|
||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||
for i, network := range vpnNetworks {
|
||||
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
||||
f.l.WithField("vpnNetworks", vpnNetworks).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Refusing to handshake with myself")
|
||||
return
|
||||
}
|
||||
vpnAddrs[i] = network.Addr()
|
||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||
anyVpnAddrsInCommon = true
|
||||
}
|
||||
}
|
||||
|
||||
if !via.IsRelayed {
|
||||
// We only want to apply the remote allow list for direct tunnels here
|
||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
||||
f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
Debug("lighthouse.remote_allow_list denied incoming handshake")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
myIndex, err := generateIndex(f.l)
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Failed to generate index")
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo := &HostInfo{
|
||||
ConnectionState: ci,
|
||||
localIndexId: myIndex,
|
||||
remoteIndexId: hs.Details.InitiatorIndex,
|
||||
vpnAddrs: vpnAddrs,
|
||||
HandshakePacket: make(map[uint8][]byte, 0),
|
||||
lastHandshakeTime: hs.Details.Time,
|
||||
relayState: RelayState{
|
||||
relays: nil,
|
||||
relayForByAddr: map[netip.Addr]*Relay{},
|
||||
relayForByIdx: map[uint32]*Relay{},
|
||||
},
|
||||
}
|
||||
|
||||
msgRxL := f.l.WithFields(m{
|
||||
"vpnAddrs": vpnAddrs,
|
||||
"from": via,
|
||||
"certName": certName,
|
||||
"certVersion": certVersion,
|
||||
"fingerprint": fingerprint,
|
||||
"issuer": issuer,
|
||||
"initiatorIndex": hs.Details.InitiatorIndex,
|
||||
"responderIndex": hs.Details.ResponderIndex,
|
||||
"remoteIndex": h.RemoteIndex,
|
||||
"handshake": m{"stage": 1, "style": "ix_psk0"},
|
||||
})
|
||||
|
||||
if anyVpnAddrsInCommon {
|
||||
msgRxL.Info("Handshake message received")
|
||||
} else {
|
||||
//todo warn if not lighthouse or relay?
|
||||
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
||||
}
|
||||
|
||||
hs.Details.ResponderIndex = myIndex
|
||||
hs.Details.Cert = cs.getHandshakeBytes(ci.myCert.Version())
|
||||
if hs.Details.Cert == nil {
|
||||
msgRxL.WithField("myCertVersion", ci.myCert.Version()).
|
||||
Error("Unable to handshake with host because no certificate handshake bytes is available")
|
||||
return
|
||||
}
|
||||
|
||||
hs.Details.CertVersion = uint32(ci.myCert.Version())
|
||||
// Update the time in case their clock is way off from ours
|
||||
hs.Details.Time = uint64(time.Now().UnixNano())
|
||||
|
||||
hsBytes, err := hs.Marshal()
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Failed to marshal handshake message")
|
||||
return
|
||||
}
|
||||
|
||||
nh := header.Encode(make([]byte, header.Len), header.Version, header.Handshake, header.HandshakeIXPSK0, hs.Details.InitiatorIndex, 2)
|
||||
msg, dKey, eKey, err := ci.H.WriteMessage(nh, hsBytes)
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Failed to call noise.WriteMessage")
|
||||
return
|
||||
} else if dKey == nil || eKey == nil {
|
||||
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).Error("Noise did not arrive at a key")
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo.HandshakePacket[0] = make([]byte, len(packet[header.Len:]))
|
||||
copy(hostinfo.HandshakePacket[0], packet[header.Len:])
|
||||
|
||||
// Regardless of whether you are the sender or receiver, you should arrive here
|
||||
// and complete standing up the connection.
|
||||
hostinfo.HandshakePacket[2] = make([]byte, len(msg))
|
||||
copy(hostinfo.HandshakePacket[2], msg)
|
||||
|
||||
// We are sending handshake packet 2, so we don't expect to receive
|
||||
// handshake packet 2 from the initiator.
|
||||
ci.window.Update(f.l, 2)
|
||||
|
||||
ci.peerCert = remoteCert
|
||||
ci.dKey = NewNebulaCipherState(dKey)
|
||||
ci.eKey = NewNebulaCipherState(eKey)
|
||||
|
||||
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
||||
if !via.IsRelayed {
|
||||
hostinfo.SetRemote(via.UdpAddr)
|
||||
}
|
||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||
|
||||
existing, err := f.handshakeManager.CheckAndComplete(hostinfo, 0, f)
|
||||
if err != nil {
|
||||
switch err {
|
||||
case ErrAlreadySeen:
|
||||
// Update remote if preferred
|
||||
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
||||
// Send a test packet to ensure the other side has also switched to
|
||||
// the preferred remote
|
||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||
}
|
||||
|
||||
msg = existing.HandshakePacket[2]
|
||||
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||
if !via.IsRelayed {
|
||||
err := f.outside.WriteTo(msg, via.UdpAddr)
|
||||
if err != nil {
|
||||
f.l.WithField("vpnAddrs", existing.vpnAddrs).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("cached", true).
|
||||
WithError(err).Error("Failed to send handshake message")
|
||||
} else {
|
||||
f.l.WithField("vpnAddrs", existing.vpnAddrs).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("cached", true).
|
||||
Info("Handshake message sent")
|
||||
}
|
||||
return
|
||||
} else {
|
||||
if via.relay == nil {
|
||||
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||
return
|
||||
}
|
||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||
f.l.WithField("vpnAddrs", existing.vpnAddrs).WithField("relay", via.relayHI.vpnAddrs[0]).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("cached", true).
|
||||
Info("Handshake message sent")
|
||||
return
|
||||
}
|
||||
case ErrExistingHostInfo:
|
||||
// This means there was an existing tunnel and this handshake was older than the one we are currently based on
|
||||
f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("oldHandshakeTime", existing.lastHandshakeTime).
|
||||
WithField("newHandshakeTime", hostinfo.lastHandshakeTime).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Info("Handshake too old")
|
||||
|
||||
// Send a test packet to trigger an authenticated tunnel test, this should suss out any lingering tunnel issues
|
||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||
return
|
||||
case ErrLocalIndexCollision:
|
||||
// This means we failed to insert because of collision on localIndexId. Just let the next handshake packet retry
|
||||
f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
WithField("localIndex", hostinfo.localIndexId).WithField("collision", existing.vpnAddrs).
|
||||
Error("Failed to add HostInfo due to localIndex collision")
|
||||
return
|
||||
default:
|
||||
// Shouldn't happen, but just in case someone adds a new error type to CheckAndComplete
|
||||
// And we forget to update it here
|
||||
f.l.WithError(err).WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 1, "style": "ix_psk0"}).
|
||||
Error("Failed to add HostInfo to HostMap")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Do the send
|
||||
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
||||
if !via.IsRelayed {
|
||||
err = f.outside.WriteTo(msg, via.UdpAddr)
|
||||
log := f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 2, "style": "ix_psk0"})
|
||||
if err != nil {
|
||||
log.WithError(err).Error("Failed to send handshake")
|
||||
} else {
|
||||
log.Info("Handshake message sent")
|
||||
}
|
||||
} else {
|
||||
if via.relay == nil {
|
||||
f.l.Error("Handshake send failed: both addr and via.relay are nil.")
|
||||
return
|
||||
}
|
||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||
// I successfully received a handshake. Just in case I marked this tunnel as 'Disestablished', ensure
|
||||
// it's correctly marked as working.
|
||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
||||
f.l.WithField("vpnAddrs", vpnAddrs).WithField("relay", via.relayHI.vpnAddrs[0]).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
Info("Handshake message sent")
|
||||
}
|
||||
|
||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||
|
||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func ixHandshakeStage2(f *Interface, via ViaSender, hh *HandshakeHostInfo, packet []byte, h *header.H) bool {
|
||||
if hh == nil {
|
||||
// Nothing here to tear down, got a bogus stage 2 packet
|
||||
return true
|
||||
}
|
||||
|
||||
hh.Lock()
|
||||
defer hh.Unlock()
|
||||
|
||||
hostinfo := hh.hostinfo
|
||||
if !via.IsRelayed {
|
||||
// The vpnAddr we know about is the one we tried to handshake with, use it to apply the remote allow list.
|
||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).Debug("lighthouse.remote_allow_list denied incoming handshake")
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
ci := hostinfo.ConnectionState
|
||||
msg, eKey, dKey, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).WithField("header", h).
|
||||
Error("Failed to call noise.ReadMessage")
|
||||
|
||||
// We don't want to tear down the connection on a bad ReadMessage because it could be an attacker trying
|
||||
// to DOS us. Every other error condition after should to allow a possible good handshake to complete in the
|
||||
// near future
|
||||
return false
|
||||
} else if dKey == nil || eKey == nil {
|
||||
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
Error("Noise did not arrive at a key")
|
||||
|
||||
// This should be impossible in IX but just in case, if we get here then there is no chance to recover
|
||||
// the handshake state machine. Tear it down
|
||||
return true
|
||||
}
|
||||
|
||||
hs := &NebulaHandshake{}
|
||||
err = hs.Unmarshal(msg)
|
||||
if err != nil || hs.Details == nil {
|
||||
f.l.WithError(err).WithField("vpnAddrs", hostinfo.vpnAddrs).WithField("from", via).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).Error("Failed unmarshal handshake message")
|
||||
|
||||
// The handshake state machine is complete, if things break now there is no chance to recover. Tear down and start again
|
||||
return true
|
||||
}
|
||||
|
||||
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||
if err != nil {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
Info("Handshake did not contain a certificate")
|
||||
return true
|
||||
}
|
||||
|
||||
remoteCert, err := f.pki.GetCAPool().VerifyCertificate(time.Now(), rc)
|
||||
if err != nil {
|
||||
fp, err := rc.Fingerprint()
|
||||
if err != nil {
|
||||
fp = "<error generating certificate fingerprint>"
|
||||
}
|
||||
|
||||
e := f.l.WithError(err).WithField("from", via).
|
||||
WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
WithField("certFingerprint", fp).
|
||||
WithField("certVpnNetworks", rc.Networks())
|
||||
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
e = e.WithField("cert", rc)
|
||||
}
|
||||
|
||||
e.Info("Invalid certificate from host")
|
||||
return true
|
||||
}
|
||||
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||
f.l.WithField("from", via).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
WithField("cert", remoteCert).Info("public key mismatch between certificate and handshake")
|
||||
return true
|
||||
}
|
||||
|
||||
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||
f.l.WithError(err).WithField("from", via).
|
||||
WithField("vpnAddrs", hostinfo.vpnAddrs).
|
||||
WithField("cert", remoteCert).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
Info("No networks in certificate")
|
||||
return true
|
||||
}
|
||||
|
||||
vpnNetworks := remoteCert.Certificate.Networks()
|
||||
certName := remoteCert.Certificate.Name()
|
||||
certVersion := remoteCert.Certificate.Version()
|
||||
fingerprint := remoteCert.Fingerprint
|
||||
issuer := remoteCert.Certificate.Issuer()
|
||||
|
||||
hostinfo.remoteIndexId = hs.Details.ResponderIndex
|
||||
hostinfo.lastHandshakeTime = hs.Details.Time
|
||||
|
||||
// Store their cert and our symmetric keys
|
||||
ci.peerCert = remoteCert
|
||||
ci.dKey = NewNebulaCipherState(dKey)
|
||||
ci.eKey = NewNebulaCipherState(eKey)
|
||||
|
||||
// Make sure the current udpAddr being used is set for responding
|
||||
if !via.IsRelayed {
|
||||
hostinfo.SetRemote(via.UdpAddr)
|
||||
} else {
|
||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
||||
}
|
||||
|
||||
correctHostResponded := false
|
||||
anyVpnAddrsInCommon := false
|
||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
||||
for i, network := range vpnNetworks {
|
||||
vpnAddrs[i] = network.Addr()
|
||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
||||
anyVpnAddrsInCommon = true
|
||||
}
|
||||
if hostinfo.vpnAddrs[0] == network.Addr() {
|
||||
// todo is it more correct to see if any of hostinfo.vpnAddrs are in the cert? it should have len==1, but one day it might not?
|
||||
correctHostResponded = true
|
||||
}
|
||||
}
|
||||
|
||||
// Ensure the right host responded
|
||||
if !correctHostResponded {
|
||||
f.l.WithField("intendedVpnAddrs", hostinfo.vpnAddrs).WithField("haveVpnNetworks", vpnNetworks).
|
||||
WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
Info("Incorrect host responded to handshake")
|
||||
|
||||
// Release our old handshake from pending, it should not continue
|
||||
f.handshakeManager.DeleteHostInfo(hostinfo)
|
||||
|
||||
// Create a new hostinfo/handshake for the intended vpn ip
|
||||
//TODO is hostinfo.vpnAddrs[0] always the address to use?
|
||||
f.handshakeManager.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
||||
// Block the current used address
|
||||
newHH.hostinfo.remotes = hostinfo.remotes
|
||||
newHH.hostinfo.remotes.BlockRemote(via)
|
||||
|
||||
f.l.WithField("blockedUdpAddrs", newHH.hostinfo.remotes.CopyBlockedRemotes()).
|
||||
WithField("vpnNetworks", vpnNetworks).
|
||||
WithField("remotes", newHH.hostinfo.remotes.CopyAddrs(f.hostMap.GetPreferredRanges())).
|
||||
Info("Blocked addresses for handshakes")
|
||||
|
||||
// Swap the packet store to benefit the original intended recipient
|
||||
newHH.packetStore = hh.packetStore
|
||||
hh.packetStore = []*cachedPacket{}
|
||||
|
||||
// Finally, put the correct vpn addrs in the host info, tell them to close the tunnel, and return true to tear down
|
||||
hostinfo.vpnAddrs = vpnAddrs
|
||||
f.sendCloseTunnel(hostinfo)
|
||||
})
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// Mark packet 2 as seen so it doesn't show up as missed
|
||||
ci.window.Update(f.l, 2)
|
||||
|
||||
duration := time.Since(hh.startTime).Nanoseconds()
|
||||
msgRxL := f.l.WithField("vpnAddrs", vpnAddrs).WithField("from", via).
|
||||
WithField("certName", certName).
|
||||
WithField("certVersion", certVersion).
|
||||
WithField("fingerprint", fingerprint).
|
||||
WithField("issuer", issuer).
|
||||
WithField("initiatorIndex", hs.Details.InitiatorIndex).WithField("responderIndex", hs.Details.ResponderIndex).
|
||||
WithField("remoteIndex", h.RemoteIndex).WithField("handshake", m{"stage": 2, "style": "ix_psk0"}).
|
||||
WithField("durationNs", duration).
|
||||
WithField("sentCachedPackets", len(hh.packetStore))
|
||||
if anyVpnAddrsInCommon {
|
||||
msgRxL.Info("Handshake message received")
|
||||
} else {
|
||||
//todo warn if not lighthouse or relay?
|
||||
msgRxL.Info("Handshake message received, but no vpnNetworks in common.")
|
||||
}
|
||||
|
||||
// Build up the radix for the firewall if we have subnets in the cert
|
||||
hostinfo.vpnAddrs = vpnAddrs
|
||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
||||
|
||||
// Complete our handshake and update metrics, this will replace any existing tunnels for the vpnAddrs here
|
||||
f.handshakeManager.Complete(hostinfo, f)
|
||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(f.l).Debugf("Sending %d stored packets", len(hh.packetStore))
|
||||
}
|
||||
|
||||
if len(hh.packetStore) > 0 {
|
||||
nb := make([]byte, 12, 12)
|
||||
out := make([]byte, mtu)
|
||||
for _, cp := range hh.packetStore {
|
||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||
}
|
||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||
}
|
||||
|
||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||
f.metricHandshakes.Update(duration)
|
||||
|
||||
return false
|
||||
}
|
||||
+655
-216
File diff suppressed because it is too large
Load Diff
+136
-1
@@ -5,6 +5,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/test"
|
||||
@@ -27,7 +28,7 @@ func Test_NewHandshakeManagerVpnIp(t *testing.T) {
|
||||
initiatingVersion: cert.Version1,
|
||||
privateKey: []byte{},
|
||||
v1Cert: &dummyCert{version: cert.Version1},
|
||||
v1HandshakeBytes: []byte{},
|
||||
v1Credential: nil,
|
||||
}
|
||||
|
||||
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
||||
@@ -100,3 +101,137 @@ func (mw *mockEncWriter) GetHostInfo(_ netip.Addr) *HostInfo {
|
||||
func (mw *mockEncWriter) GetCertState() *CertState {
|
||||
return &CertState{initiatingVersion: cert.Version2}
|
||||
}
|
||||
|
||||
func TestValidatePeerCert(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
|
||||
myNetwork := netip.MustParsePrefix("10.0.0.1/24")
|
||||
myAddrTable := new(bart.Lite)
|
||||
myAddrTable.Insert(netip.PrefixFrom(myNetwork.Addr(), myNetwork.Addr().BitLen()))
|
||||
myNetTable := new(bart.Lite)
|
||||
myNetTable.Insert(myNetwork.Masked())
|
||||
|
||||
newHM := func() *HandshakeManager {
|
||||
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
||||
hm.f = &Interface{
|
||||
handshakeManager: hm,
|
||||
pki: &PKI{},
|
||||
l: l,
|
||||
myVpnAddrsTable: myAddrTable,
|
||||
myVpnNetworksTable: myNetTable,
|
||||
lightHouse: hm.lightHouse,
|
||||
}
|
||||
return hm
|
||||
}
|
||||
|
||||
cached := func(networks ...netip.Prefix) *cert.CachedCertificate {
|
||||
return &cert.CachedCertificate{
|
||||
Certificate: &dummyCert{name: "peer", networks: networks},
|
||||
}
|
||||
}
|
||||
|
||||
via := ViaSender{
|
||||
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
||||
IsRelayed: true, // skip the remote allow list (covered separately)
|
||||
}
|
||||
|
||||
t.Run("addr inside our networks sets anyVpnAddrsInCommon", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
// 10.0.0.2 falls inside our 10.0.0.0/24
|
||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.2/24")))
|
||||
assert.True(t, ok)
|
||||
assert.True(t, common)
|
||||
assert.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.2")}, addrs)
|
||||
})
|
||||
|
||||
t.Run("addr outside our networks leaves anyVpnAddrsInCommon false", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("192.168.1.5/24")))
|
||||
assert.True(t, ok)
|
||||
assert.False(t, common)
|
||||
assert.Equal(t, []netip.Addr{netip.MustParseAddr("192.168.1.5")}, addrs)
|
||||
})
|
||||
|
||||
t.Run("any matching network is enough", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
addrs, common, ok := hm.validatePeerCert(via, cached(
|
||||
netip.MustParsePrefix("192.168.1.5/24"),
|
||||
netip.MustParsePrefix("10.0.0.42/24"),
|
||||
))
|
||||
assert.True(t, ok)
|
||||
assert.True(t, common)
|
||||
assert.Len(t, addrs, 2)
|
||||
})
|
||||
|
||||
t.Run("self-handshake is rejected", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
// 10.0.0.1 is in myVpnAddrsTable
|
||||
addrs, common, ok := hm.validatePeerCert(via, cached(netip.MustParsePrefix("10.0.0.1/24")))
|
||||
assert.False(t, ok)
|
||||
assert.False(t, common)
|
||||
assert.Nil(t, addrs)
|
||||
})
|
||||
|
||||
t.Run("cert with no networks is rejected", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
addrs, common, ok := hm.validatePeerCert(via, cached())
|
||||
assert.False(t, ok)
|
||||
assert.False(t, common)
|
||||
assert.Nil(t, addrs)
|
||||
})
|
||||
}
|
||||
|
||||
func TestHandleIncomingDispatch(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
|
||||
newHM := func() *HandshakeManager {
|
||||
hm := NewHandshakeManager(l, newHostMap(l), newTestLighthouse(), &udp.NoopConn{}, defaultHandshakeConfig)
|
||||
hm.f = &Interface{
|
||||
handshakeManager: hm,
|
||||
pki: &PKI{},
|
||||
l: l,
|
||||
}
|
||||
return hm
|
||||
}
|
||||
|
||||
via := ViaSender{
|
||||
UdpAddr: netip.MustParseAddrPort("198.51.100.7:4242"),
|
||||
IsRelayed: true, // bypass remote allow list
|
||||
}
|
||||
|
||||
// A packet body of zero length is fine for these tests: dispatch is
|
||||
// gated on header fields, and we assert that we never reach noise/cert
|
||||
// processing for any of the malformed shapes here.
|
||||
pkt := make([]byte, header.Len)
|
||||
|
||||
t.Run("unsupported subtype dropped", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
h := &header.H{Type: header.Handshake, Subtype: header.MessageSubType(99), MessageCounter: 1}
|
||||
hm.HandleIncoming(via, pkt, h)
|
||||
assert.Empty(t, hm.indexes, "no pending handshake should be created")
|
||||
})
|
||||
|
||||
t.Run("stage-1 with non-zero RemoteIndex dropped", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
h := &header.H{
|
||||
Type: header.Handshake,
|
||||
Subtype: header.HandshakeIXPSK0,
|
||||
RemoteIndex: 0xdeadbeef,
|
||||
MessageCounter: 1,
|
||||
}
|
||||
hm.HandleIncoming(via, pkt, h)
|
||||
assert.Empty(t, hm.indexes, "spoofed stage-1 must not create a pending machine")
|
||||
})
|
||||
|
||||
t.Run("continuation with no matching pending index dropped", func(t *testing.T) {
|
||||
hm := newHM()
|
||||
h := &header.H{
|
||||
Type: header.Handshake,
|
||||
Subtype: header.HandshakeIXPSK0,
|
||||
RemoteIndex: 0xcafef00d,
|
||||
MessageCounter: 2,
|
||||
}
|
||||
hm.HandleIncoming(via, pkt, h)
|
||||
assert.Empty(t, hm.indexes, "orphan stage-2 must not create state")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -174,6 +174,10 @@ func (h *H) SubTypeName() string {
|
||||
return SubTypeName(h.Type, h.Subtype)
|
||||
}
|
||||
|
||||
func (h *H) IsValidSubType() bool {
|
||||
return IsValidSubType(h.Type, h.Subtype)
|
||||
}
|
||||
|
||||
// SubTypeName will transform a nebula message sub type into a human string
|
||||
func SubTypeName(t MessageType, s MessageSubType) string {
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
@@ -185,6 +189,16 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
||||
return "unknown"
|
||||
}
|
||||
|
||||
func IsValidSubType(t MessageType, s MessageSubType) bool {
|
||||
if n, ok := subTypeMap[t]; ok {
|
||||
if _, ok := (*n)[s]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// NewHeader turns bytes into a header
|
||||
func NewHeader(b []byte) (*H, error) {
|
||||
h := new(H)
|
||||
|
||||
+49
-31
@@ -1,9 +1,11 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
@@ -13,10 +15,10 @@ import (
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
)
|
||||
|
||||
const defaultPromoteEvery = 1000 // Count of packets sent before we try moving a tunnel to a preferred underlay ip address
|
||||
@@ -60,7 +62,7 @@ type HostMap struct {
|
||||
RemoteIndexes map[uint32]*HostInfo
|
||||
Hosts map[netip.Addr]*HostInfo
|
||||
preferredRanges atomic.Pointer[[]netip.Prefix]
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// For synchronization, treat the pointed-to Relay struct as immutable. To edit the Relay
|
||||
@@ -313,7 +315,7 @@ type cachedPacketMetrics struct {
|
||||
dropped metrics.Counter
|
||||
}
|
||||
|
||||
func NewHostMapFromConfig(l *logrus.Logger, c *config.C) *HostMap {
|
||||
func NewHostMapFromConfig(l *slog.Logger, c *config.C) *HostMap {
|
||||
hm := newHostMap(l)
|
||||
|
||||
hm.reload(c, true)
|
||||
@@ -321,13 +323,12 @@ func NewHostMapFromConfig(l *logrus.Logger, c *config.C) *HostMap {
|
||||
hm.reload(c, false)
|
||||
})
|
||||
|
||||
l.WithField("preferredRanges", hm.GetPreferredRanges()).
|
||||
Info("Main HostMap created")
|
||||
l.Info("Main HostMap created", "preferredRanges", hm.GetPreferredRanges())
|
||||
|
||||
return hm
|
||||
}
|
||||
|
||||
func newHostMap(l *logrus.Logger) *HostMap {
|
||||
func newHostMap(l *slog.Logger) *HostMap {
|
||||
return &HostMap{
|
||||
Indexes: map[uint32]*HostInfo{},
|
||||
Relays: map[uint32]*HostInfo{},
|
||||
@@ -346,7 +347,10 @@ func (hm *HostMap) reload(c *config.C, initial bool) {
|
||||
preferredRange, err := netip.ParsePrefix(rawPreferredRange)
|
||||
|
||||
if err != nil {
|
||||
hm.l.WithError(err).WithField("range", rawPreferredRanges).Warn("Failed to parse preferred ranges, ignoring")
|
||||
hm.l.Warn("Failed to parse preferred ranges, ignoring",
|
||||
"error", err,
|
||||
"range", rawPreferredRanges,
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -355,7 +359,10 @@ func (hm *HostMap) reload(c *config.C, initial bool) {
|
||||
|
||||
oldRanges := hm.preferredRanges.Swap(&preferredRanges)
|
||||
if !initial {
|
||||
hm.l.WithField("oldPreferredRanges", *oldRanges).WithField("newPreferredRanges", preferredRanges).Info("preferred_ranges changed")
|
||||
hm.l.Info("preferred_ranges changed",
|
||||
"oldPreferredRanges", *oldRanges,
|
||||
"newPreferredRanges", preferredRanges,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -488,10 +495,11 @@ func (hm *HostMap) unlockedInnerDeleteHostInfo(hostinfo *HostInfo, addr netip.Ad
|
||||
hm.Indexes = map[uint32]*HostInfo{}
|
||||
}
|
||||
|
||||
if hm.l.Level >= logrus.DebugLevel {
|
||||
hm.l.WithField("hostMap", m{"mapTotalSize": len(hm.Hosts),
|
||||
"vpnAddrs": hostinfo.vpnAddrs, "indexNumber": hostinfo.localIndexId, "remoteIndexNumber": hostinfo.remoteIndexId}).
|
||||
Debug("Hostmap hostInfo deleted")
|
||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hm.l.Debug("Hostmap hostInfo deleted",
|
||||
"hostMap", m{"mapTotalSize": len(hm.Hosts),
|
||||
"vpnAddrs": hostinfo.vpnAddrs, "indexNumber": hostinfo.localIndexId, "remoteIndexNumber": hostinfo.remoteIndexId},
|
||||
)
|
||||
}
|
||||
|
||||
if isLastHostinfo {
|
||||
@@ -604,9 +612,9 @@ func (hm *HostMap) queryVpnAddr(vpnIp netip.Addr, promoteIfce *Interface) *HostI
|
||||
// unlockedAddHostInfo assumes you have a write-lock and will add a hostinfo object to the hostmap Indexes and RemoteIndexes maps.
|
||||
// If an entry exists for the Hosts table (vpnIp -> hostinfo) then the provided hostinfo will be made primary
|
||||
func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
||||
if f.serveDns {
|
||||
if f.dnsServer != nil {
|
||||
remoteCert := hostinfo.ConnectionState.peerCert
|
||||
dnsR.Add(remoteCert.Certificate.Name()+".", hostinfo.vpnAddrs)
|
||||
f.dnsServer.Add(remoteCert.Certificate.Name()+".", hostinfo.vpnAddrs)
|
||||
}
|
||||
for _, addr := range hostinfo.vpnAddrs {
|
||||
hm.unlockedInnerAddHostInfo(addr, hostinfo, f)
|
||||
@@ -615,10 +623,11 @@ func (hm *HostMap) unlockedAddHostInfo(hostinfo *HostInfo, f *Interface) {
|
||||
hm.Indexes[hostinfo.localIndexId] = hostinfo
|
||||
hm.RemoteIndexes[hostinfo.remoteIndexId] = hostinfo
|
||||
|
||||
if hm.l.Level >= logrus.DebugLevel {
|
||||
hm.l.WithField("hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
||||
"hostinfo": m{"existing": true, "localIndexId": hostinfo.localIndexId, "vpnAddrs": hostinfo.vpnAddrs}}).
|
||||
Debug("Hostmap vpnIp added")
|
||||
if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hm.l.Debug("Hostmap vpnIp added",
|
||||
"hostMap", m{"vpnAddrs": hostinfo.vpnAddrs, "mapTotalSize": len(hm.Hosts),
|
||||
"hostinfo": m{"existing": true, "localIndexId": hostinfo.localIndexId, "vpnAddrs": hostinfo.vpnAddrs}},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -784,18 +793,21 @@ func (i *HostInfo) buildNetworks(myVpnNetworksTable *bart.Lite, c cert.Certifica
|
||||
}
|
||||
}
|
||||
|
||||
func (i *HostInfo) logger(l *logrus.Logger) *logrus.Entry {
|
||||
// logger returns a derived slog.Logger with per-hostinfo fields pre-bound.
|
||||
func (i *HostInfo) logger(l *slog.Logger) *slog.Logger {
|
||||
if i == nil {
|
||||
return logrus.NewEntry(l)
|
||||
return l
|
||||
}
|
||||
|
||||
li := l.WithField("vpnAddrs", i.vpnAddrs).
|
||||
WithField("localIndex", i.localIndexId).
|
||||
WithField("remoteIndex", i.remoteIndexId)
|
||||
li := l.With(
|
||||
"vpnAddrs", i.vpnAddrs,
|
||||
"localIndex", i.localIndexId,
|
||||
"remoteIndex", i.remoteIndexId,
|
||||
)
|
||||
|
||||
if connState := i.ConnectionState; connState != nil {
|
||||
if peerCert := connState.peerCert; peerCert != nil {
|
||||
li = li.WithField("certName", peerCert.Certificate.Name())
|
||||
li = li.With("certName", peerCert.Certificate.Name())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -804,14 +816,17 @@ func (i *HostInfo) logger(l *logrus.Logger) *logrus.Entry {
|
||||
|
||||
// Utility functions
|
||||
|
||||
func localAddrs(l *logrus.Logger, allowList *LocalAllowList) []netip.Addr {
|
||||
func localAddrs(l *slog.Logger, allowList *LocalAllowList) []netip.Addr {
|
||||
//FIXME: This function is pretty garbage
|
||||
var finalAddrs []netip.Addr
|
||||
ifaces, _ := net.Interfaces()
|
||||
for _, i := range ifaces {
|
||||
allow := allowList.AllowName(i.Name)
|
||||
if l.Level >= logrus.TraceLevel {
|
||||
l.WithField("interfaceName", i.Name).WithField("allow", allow).Trace("localAllowList.AllowName")
|
||||
if l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
l.Log(context.Background(), logging.LevelTrace, "localAllowList.AllowName",
|
||||
"interfaceName", i.Name,
|
||||
"allow", allow,
|
||||
)
|
||||
}
|
||||
|
||||
if !allow {
|
||||
@@ -829,8 +844,8 @@ func localAddrs(l *logrus.Logger, allowList *LocalAllowList) []netip.Addr {
|
||||
}
|
||||
|
||||
if !addr.IsValid() {
|
||||
if l.Level >= logrus.DebugLevel {
|
||||
l.WithField("localAddr", rawAddr).Debug("addr was invalid")
|
||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
l.Debug("addr was invalid", "localAddr", rawAddr)
|
||||
}
|
||||
continue
|
||||
}
|
||||
@@ -838,8 +853,11 @@ func localAddrs(l *logrus.Logger, allowList *LocalAllowList) []netip.Addr {
|
||||
|
||||
if addr.IsLoopback() == false && addr.IsLinkLocalUnicast() == false {
|
||||
isAllowed := allowList.Allow(addr)
|
||||
if l.Level >= logrus.TraceLevel {
|
||||
l.WithField("localAddr", addr).WithField("allowed", isAllowed).Trace("localAllowList.Allow")
|
||||
if l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
l.Log(context.Background(), logging.LevelTrace, "localAllowList.Allow",
|
||||
"localAddr", addr,
|
||||
"allowed", isAllowed,
|
||||
)
|
||||
}
|
||||
if !isAllowed {
|
||||
continue
|
||||
|
||||
+1
-1
@@ -196,7 +196,7 @@ func TestHostMap_DeleteHostInfo(t *testing.T) {
|
||||
|
||||
func TestHostMap_reload(t *testing.T) {
|
||||
l := test.NewLogger()
|
||||
c := config.NewC(l)
|
||||
c := config.NewC(test.NewLogger())
|
||||
|
||||
hm := NewHostMapFromConfig(l, c)
|
||||
|
||||
|
||||
@@ -1,25 +1,46 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/iputil"
|
||||
"github.com/slackhq/nebula/noiseutil"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/routing"
|
||||
)
|
||||
|
||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
||||
err := newPacket(packet, false, fwPacket)
|
||||
if err != nil {
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("packet", packet).Debugf("Error while validating outbound packet: %s", err)
|
||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||
// 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) {
|
||||
f.l.Debug("Error while validating outbound packet",
|
||||
"packet", packet,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
parsed.Key.Hydrate(fwPacket)
|
||||
|
||||
// Ignore local broadcast packets
|
||||
if f.dropLocalBroadcast {
|
||||
if f.myBroadcastAddrsTable.Contains(fwPacket.RemoteAddr) {
|
||||
@@ -33,9 +54,16 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||
// TUN device.
|
||||
if immediatelyForwardToSelf {
|
||||
_, err := f.readers[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 {
|
||||
f.l.WithError(err).Error("Failed to forward to tun")
|
||||
f.l.Error("Failed to forward to tun", "error", err)
|
||||
}
|
||||
}
|
||||
// Otherwise, drop. On linux, we should never see these packets - Linux
|
||||
@@ -49,15 +77,28 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
}
|
||||
|
||||
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 {
|
||||
f.rejectInside(packet, out, q)
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("vpnAddr", fwPacket.RemoteAddr).
|
||||
WithField("fwPacket", fwPacket).
|
||||
Debugln("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks")
|
||||
f.rejectInside(packet, rejectBuf, q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||
"vpnAddr", fwPacket.RemoteAddr,
|
||||
"fwPacket", fwPacket,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -66,21 +107,163 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
||||
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 {
|
||||
f.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
||||
|
||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
||||
} else {
|
||||
f.rejectInside(packet, out, q)
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(f.l).
|
||||
WithField("fwPacket", fwPacket).
|
||||
WithField("reason", dropReason).
|
||||
Debugln("dropping outbound packet")
|
||||
f.rejectInside(packet, rejectBuf, q)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||
"fwPacket", fwPacket,
|
||||
"reason", dropReason,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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) {
|
||||
if !f.firewall.InSendReject {
|
||||
return
|
||||
@@ -93,7 +276,7 @@ func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||
|
||||
_, err := f.readers[q].Write(out)
|
||||
if err != nil {
|
||||
f.l.WithError(err).Error("Failed to write to tun")
|
||||
f.l.Error("Failed to write to tun", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -108,11 +291,11 @@ func (f *Interface) rejectOutside(packet []byte, ci *ConnectionState, hostinfo *
|
||||
}
|
||||
|
||||
if len(out) > iputil.MaxRejectPacketSize {
|
||||
if f.l.GetLevel() >= logrus.InfoLevel {
|
||||
f.l.
|
||||
WithField("packet", packet).
|
||||
WithField("outPacket", out).
|
||||
Info("rejectOutside: packet too big, not sending")
|
||||
if f.l.Enabled(context.Background(), slog.LevelInfo) {
|
||||
f.l.Info("rejectOutside: packet too big, not sending",
|
||||
"packet", packet,
|
||||
"outPacket", out,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -184,10 +367,11 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
||||
// This would also need to interact with unsafe_route updates through reloading the config or
|
||||
// use of the use_system_route_table option
|
||||
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("destination", destinationAddr).
|
||||
WithField("originalGateway", gatewayAddr).
|
||||
Debugln("Calculated gateway for ECMP not available, attempting other gateways")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Calculated gateway for ECMP not available, attempting other gateways",
|
||||
"destination", destinationAddr,
|
||||
"originalGateway", gatewayAddr,
|
||||
)
|
||||
}
|
||||
|
||||
for i := range gateways {
|
||||
@@ -210,20 +394,22 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
||||
}
|
||||
|
||||
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
||||
fp := &firewall.Packet{}
|
||||
err := newPacket(p, false, fp)
|
||||
if err != nil {
|
||||
f.l.Warnf("error while parsing outgoing packet for firewall check; %v", err)
|
||||
var parsed batch.RxParsed
|
||||
if err := batch.ParsePacket(p, false, &parsed); err != nil {
|
||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
||||
return
|
||||
}
|
||||
fp := &firewall.Packet{}
|
||||
parsed.Key.Hydrate(fp)
|
||||
|
||||
// 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 f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("fwPacket", fp).
|
||||
WithField("reason", dropReason).
|
||||
Debugln("dropping cached packet")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping cached packet",
|
||||
"fwPacket", fp,
|
||||
"reason", dropReason,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -239,9 +425,10 @@ func (f *Interface) SendMessageToVpnAddr(t header.MessageType, st header.Message
|
||||
})
|
||||
|
||||
if hostInfo == nil {
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("vpnAddr", vpnAddr).
|
||||
Debugln("dropping SendMessageToVpnAddr, vpnAddr not in our vpn networks or in unsafe routes")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("dropping SendMessageToVpnAddr, vpnAddr not in our vpn networks or in unsafe routes",
|
||||
"vpnAddr", vpnAddr,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -267,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)
|
||||
}
|
||||
|
||||
// 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,
|
||||
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||
relay *Relay,
|
||||
ad,
|
||||
nb,
|
||||
out []byte,
|
||||
nocopy bool,
|
||||
) {
|
||||
) ([]byte, error) {
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||
via.ConnectionState.writeLock.Lock()
|
||||
@@ -297,13 +476,13 @@ func (f *Interface) SendVia(via *HostInfo,
|
||||
if noiseutil.EncryptLockNeeded {
|
||||
via.ConnectionState.writeLock.Unlock()
|
||||
}
|
||||
via.logger(f.l).
|
||||
WithField("outCap", cap(out)).
|
||||
WithField("payloadLen", len(ad)).
|
||||
WithField("headerLen", len(out)).
|
||||
WithField("cipherOverhead", via.ConnectionState.eKey.Overhead()).
|
||||
Error("SendVia out buffer not large enough for relay")
|
||||
return
|
||||
via.logger(f.l).Error("SendVia out buffer not large enough for relay",
|
||||
"outCap", cap(out),
|
||||
"payloadLen", len(ad),
|
||||
"headerLen", len(out),
|
||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||
)
|
||||
return nil, io.ErrShortBuffer
|
||||
}
|
||||
|
||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||
@@ -322,14 +501,33 @@ func (f *Interface) SendVia(via *HostInfo,
|
||||
via.ConnectionState.writeLock.Unlock()
|
||||
}
|
||||
if err != nil {
|
||||
via.logger(f.l).WithError(err).Info("Failed to EncryptDanger in sendVia")
|
||||
return
|
||||
}
|
||||
err = f.writers[0].WriteTo(out, via.remote)
|
||||
if err != nil {
|
||||
via.logger(f.l).WithError(err).Info("Failed to WriteTo in sendVia")
|
||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||
return nil, err
|
||||
}
|
||||
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 {
|
||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||
@@ -366,8 +564,10 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
// 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.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("vpnAddrs", hostinfo.vpnAddrs).Debug("Lighthouse update triggered for punch due to rebind counter")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -377,24 +577,30 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
ci.writeLock.Unlock()
|
||||
}
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).
|
||||
WithField("udpAddr", remote).WithField("counter", c).
|
||||
WithField("attemptedCounter", c).
|
||||
Error("Failed to encrypt outgoing packet")
|
||||
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", remote,
|
||||
"counter", c,
|
||||
"attemptedCounter", c,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
if remote.IsValid() {
|
||||
err = f.writers[q].WriteTo(out, remote)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).
|
||||
WithField("udpAddr", remote).Error("Failed to write outgoing packet")
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", remote,
|
||||
)
|
||||
}
|
||||
} else if hostinfo.remote.IsValid() {
|
||||
err = f.writers[q].WriteTo(out, hostinfo.remote)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).
|
||||
WithField("udpAddr", remote).Error("Failed to write outgoing packet")
|
||||
hostinfo.logger(f.l).Error("Failed to write outgoing packet",
|
||||
"error", err,
|
||||
"udpAddr", remote,
|
||||
)
|
||||
}
|
||||
} else {
|
||||
// Try to send via a relay
|
||||
@@ -402,7 +608,10 @@ func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType
|
||||
relayHostInfo, relay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||
if err != nil {
|
||||
hostinfo.relayState.DeleteRelay(relayIP)
|
||||
hostinfo.logger(f.l).WithField("relay", relayIP).WithError(err).Info("sendNoMetrics failed to find HostInfo")
|
||||
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
||||
"relay", relayIP,
|
||||
"error", err,
|
||||
)
|
||||
continue
|
||||
}
|
||||
f.SendVia(relayHostInfo, relay, out, nb, fullOut[:header.Len+len(out)], true)
|
||||
|
||||
+200
-70
@@ -4,19 +4,23 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
"runtime"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/util"
|
||||
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/slackhq/nebula/overlay/tio"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
)
|
||||
|
||||
@@ -29,7 +33,7 @@ type InterfaceConfig struct {
|
||||
pki *PKI
|
||||
Cipher string
|
||||
Firewall *Firewall
|
||||
ServeDns bool
|
||||
DnsServer *dnsServer
|
||||
HandshakeManager *HandshakeManager
|
||||
lightHouse *LightHouse
|
||||
connectionManager *connectionManager
|
||||
@@ -46,7 +50,14 @@ type InterfaceConfig struct {
|
||||
reQueryWait time.Duration
|
||||
|
||||
ConntrackCacheTimeout time.Duration
|
||||
l *logrus.Logger
|
||||
|
||||
// CpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
||||
// should pin to. Queue i pins to CpuAffinity[i % len(CpuAffinity)] —
|
||||
// shorter lists than `routines` cycle. Empty list keeps the default
|
||||
// pin-to-(i % NumCPU) behavior.
|
||||
CpuAffinity []int
|
||||
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type Interface struct {
|
||||
@@ -57,7 +68,7 @@ type Interface struct {
|
||||
firewall *Firewall
|
||||
connectionManager *connectionManager
|
||||
handshakeManager *HandshakeManager
|
||||
serveDns bool
|
||||
dnsServer *dnsServer
|
||||
createTime time.Time
|
||||
lightHouse *LightHouse
|
||||
myBroadcastAddrsTable *bart.Lite
|
||||
@@ -70,7 +81,16 @@ type Interface struct {
|
||||
routines int
|
||||
disconnectInvalid atomic.Bool
|
||||
closed atomic.Bool
|
||||
relayManager *relayManager
|
||||
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
||||
// should pin to. Queue i pins to cpuAffinity[i % len(cpuAffinity)].
|
||||
// Empty falls back to the default pin-to-(i % NumCPU) behavior.
|
||||
cpuAffinity []int
|
||||
// ecnEnabled gates RFC 6040 underlay ECN propagation. When true,
|
||||
// inside.go copies the inner ECN onto the outer carrier on encap and
|
||||
// decryptToTun folds outer CE into the inner header on decap. Toggle
|
||||
// via tunnels.ecn (default true).
|
||||
ecnEnabled atomic.Bool
|
||||
relayManager *relayManager
|
||||
|
||||
tryPromoteEvery atomic.Uint32
|
||||
reQueryEvery atomic.Uint32
|
||||
@@ -85,14 +105,26 @@ type Interface struct {
|
||||
|
||||
conntrackCacheTimeout time.Duration
|
||||
|
||||
ctx context.Context
|
||||
writers []udp.Conn
|
||||
readers []io.ReadWriteCloser
|
||||
readers []tio.Queue
|
||||
// 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.
|
||||
// nil means "no fatal error" (yet)
|
||||
fatalErr atomic.Pointer[error]
|
||||
// triggerShutdown is a function that will be run exactly once, when onFatal swaps something non-nil into fatalErr
|
||||
triggerShutdown func()
|
||||
|
||||
metricHandshakes metrics.Histogram
|
||||
messageMetrics *MessageMetrics
|
||||
cachedPacketMetrics *cachedPacketMetrics
|
||||
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
type EncWriter interface {
|
||||
@@ -163,12 +195,13 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
|
||||
cs := c.pki.getCertState()
|
||||
ifce := &Interface{
|
||||
ctx: ctx,
|
||||
pki: c.pki,
|
||||
hostMap: c.HostMap,
|
||||
outside: c.Outside,
|
||||
inside: c.Inside,
|
||||
firewall: c.Firewall,
|
||||
serveDns: c.ServeDns,
|
||||
dnsServer: c.DnsServer,
|
||||
handshakeManager: c.HandshakeManager,
|
||||
createTime: time.Now(),
|
||||
lightHouse: c.lightHouse,
|
||||
@@ -177,7 +210,8 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
routines: c.routines,
|
||||
version: c.version,
|
||||
writers: make([]udp.Conn, c.routines),
|
||||
readers: make([]io.ReadWriteCloser, c.routines),
|
||||
readers: make([]tio.Queue, c.routines),
|
||||
batchers: make([]batch.RxBatcher, c.routines),
|
||||
myVpnNetworks: cs.myVpnNetworks,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
myVpnAddrs: cs.myVpnAddrs,
|
||||
@@ -186,6 +220,7 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
relayManager: c.relayManager,
|
||||
connectionManager: c.connectionManager,
|
||||
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
||||
cpuAffinity: c.CpuAffinity,
|
||||
|
||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||
messageMetrics: c.MessageMetrics,
|
||||
@@ -209,18 +244,21 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
||||
// activate creates the interface on the host. After the interface is created, any
|
||||
// other services that want to bind listeners to its IP may do so successfully. However,
|
||||
// the interface isn't going to process anything until run() is called.
|
||||
func (f *Interface) activate() {
|
||||
func (f *Interface) activate() error {
|
||||
// actually turn on tun dev
|
||||
|
||||
addr, err := f.outside.LocalAddr()
|
||||
if err != nil {
|
||||
f.l.WithError(err).Error("Failed to get udp listen address")
|
||||
f.l.Error("Failed to get udp listen address", "error", err)
|
||||
}
|
||||
|
||||
f.l.WithField("interface", f.inside.Name()).WithField("networks", f.myVpnNetworks).
|
||||
WithField("build", f.version).WithField("udpAddr", addr).
|
||||
WithField("boringcrypto", boringEnabled()).
|
||||
Info("Nebula interface is active")
|
||||
f.l.Info("Nebula interface is active",
|
||||
"interface", f.inside.Name(),
|
||||
"networks", f.myVpnNetworks,
|
||||
"build", f.version,
|
||||
"udpAddr", addr,
|
||||
"boringcrypto", boringEnabled(),
|
||||
)
|
||||
|
||||
if f.routines > 1 {
|
||||
if !f.inside.SupportsMultiqueue() || !f.outside.SupportsMultipleReaders() {
|
||||
@@ -232,32 +270,69 @@ func (f *Interface) activate() {
|
||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||
|
||||
// Prepare n tun queues
|
||||
var reader io.ReadWriteCloser = f.inside
|
||||
for i := 0; i < f.routines; i++ {
|
||||
if i > 0 {
|
||||
reader, err = f.inside.NewMultiQueueReader()
|
||||
if err != nil {
|
||||
f.l.Fatal(err)
|
||||
if err = f.inside.NewMultiQueueReader(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
f.readers[i] = reader
|
||||
}
|
||||
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])
|
||||
}
|
||||
}
|
||||
|
||||
if err := f.inside.Activate(); err != nil {
|
||||
f.wg.Add(1) // for us to wait on Close() to return
|
||||
if err = f.inside.Activate(); err != nil {
|
||||
f.wg.Done()
|
||||
f.inside.Close()
|
||||
f.l.Fatal(err)
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *Interface) run() {
|
||||
func (f *Interface) run() (func() error, error) {
|
||||
// Launch n queues to read packets from udp
|
||||
for i := 0; i < f.routines; i++ {
|
||||
go f.listenOut(i)
|
||||
f.wg.Go(func() {
|
||||
f.listenOut(i)
|
||||
})
|
||||
}
|
||||
|
||||
// Launch n queues to read packets from tun dev
|
||||
for i := 0; i < f.routines; i++ {
|
||||
go f.listenIn(f.readers[i], i)
|
||||
f.wg.Go(func() {
|
||||
f.listenIn(f.readers[i], i)
|
||||
})
|
||||
}
|
||||
|
||||
return func() error {
|
||||
f.wg.Wait()
|
||||
if e := f.fatalErr.Load(); e != nil {
|
||||
return *e
|
||||
}
|
||||
return nil
|
||||
}, nil
|
||||
}
|
||||
|
||||
// onFatal stores the first fatal reader error, and calls triggerShutdown if it was the first one
|
||||
func (f *Interface) onFatal(err error) {
|
||||
swapped := f.fatalErr.CompareAndSwap(nil, &err)
|
||||
if !swapped {
|
||||
return
|
||||
}
|
||||
if f.triggerShutdown != nil {
|
||||
f.triggerShutdown()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -269,40 +344,80 @@ func (f *Interface) listenOut(i int) {
|
||||
li = f.outside
|
||||
}
|
||||
|
||||
ctCache := firewall.NewConntrackCacheTicker(f.conntrackCacheTimeout)
|
||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
lhh := f.lightHouse.NewRequestHandler()
|
||||
plaintext := make([]byte, udp.MTU)
|
||||
h := &header.H{}
|
||||
fwPacket := &firewall.Packet{}
|
||||
parsedRx := &batch.RxParsed{}
|
||||
nb := make([]byte, 12, 12)
|
||||
|
||||
li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get(f.l))
|
||||
})
|
||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||
plaintext := f.batchers[i].Reserve(len(payload))
|
||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, parsedRx, lhh, nb, i, ctCache.Get(), meta)
|
||||
}
|
||||
|
||||
flusher := func() {
|
||||
if err := f.batchers[i].Flush(); err != nil {
|
||||
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
err := li.ListenOut(listener, flusher)
|
||||
|
||||
if err != nil && !f.closed.Load() {
|
||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||
f.onFatal(err)
|
||||
}
|
||||
|
||||
f.l.Debug("underlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
||||
packet := make([]byte, mtu)
|
||||
out := make([]byte, mtu)
|
||||
func (f *Interface) listenIn(reader tio.Queue, i int) {
|
||||
// Pin this goroutine to one CPU. LockOSThread alone keeps the goroutine
|
||||
// on a single OS thread but the kernel can still migrate that thread
|
||||
// across CPUs — XPS reads smp_processor_id() at sendmmsg time and picks
|
||||
// the TX ring from the current CPU's xps_cpus map, so an unpinned
|
||||
// thread bouncing between CPUs spreads one nebula flow's packets across
|
||||
// multiple TX rings, which the rings then drain at independent rates
|
||||
// and the wire delivers reordered.
|
||||
//
|
||||
// Pinning keeps every sendmmsg from this goroutine going through the
|
||||
// same TX ring, so the wire sees per-flow order. Cost: less scheduler
|
||||
// flexibility — if i % NumCPU collides between two TUN reader
|
||||
// goroutines they share a CPU.
|
||||
cpu := i % runtime.NumCPU()
|
||||
if n := len(f.cpuAffinity); n > 0 {
|
||||
cpu = f.cpuAffinity[i%n]
|
||||
}
|
||||
if err := util.PinThreadToCPU(cpu); err != nil {
|
||||
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{}
|
||||
nb := make([]byte, 12, 12)
|
||||
|
||||
conntrackCache := firewall.NewConntrackCacheTicker(f.conntrackCacheTimeout)
|
||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||
|
||||
for {
|
||||
n, err := reader.Read(packet)
|
||||
pkts, err := reader.Read()
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrClosed) && f.closed.Load() {
|
||||
return
|
||||
if !f.closed.Load() {
|
||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
||||
f.onFatal(err)
|
||||
}
|
||||
|
||||
f.l.WithError(err).Error("Error while reading outbound packet")
|
||||
// This only seems to happen when something fatal happens to the fd, so exit.
|
||||
os.Exit(2)
|
||||
break
|
||||
}
|
||||
|
||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get(f.l))
|
||||
for _, pkt := range pkts {
|
||||
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||
}
|
||||
if err := sb.Flush(); err != nil {
|
||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
||||
}
|
||||
}
|
||||
|
||||
f.l.Debug("overlay reader is done", "reader", i)
|
||||
}
|
||||
|
||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||
@@ -311,6 +426,7 @@ func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||
c.RegisterReloadCallback(f.reloadMisc)
|
||||
c.RegisterReloadCallback(f.reloadEcn)
|
||||
|
||||
for _, udpConn := range f.writers {
|
||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||
@@ -322,7 +438,7 @@ func (f *Interface) reloadDisconnectInvalid(c *config.C) {
|
||||
if initial || c.HasChanged("pki.disconnect_invalid") {
|
||||
f.disconnectInvalid.Store(c.GetBool("pki.disconnect_invalid", true))
|
||||
if !initial {
|
||||
f.l.Infof("pki.disconnect_invalid changed to %v", f.disconnectInvalid.Load())
|
||||
f.l.Info("pki.disconnect_invalid changed", "value", f.disconnectInvalid.Load())
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -336,7 +452,7 @@ func (f *Interface) reloadFirewall(c *config.C) {
|
||||
|
||||
fw, err := NewFirewallFromConfig(f.l, f.pki.getCertState(), c)
|
||||
if err != nil {
|
||||
f.l.WithError(err).Error("Error while creating firewall during reload")
|
||||
f.l.Error("Error while creating firewall during reload", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -349,10 +465,11 @@ func (f *Interface) reloadFirewall(c *config.C) {
|
||||
// If rulesVersion is back to zero, we have wrapped all the way around. Be
|
||||
// safe and just reset conntrack in this case.
|
||||
if fw.rulesVersion == 0 {
|
||||
f.l.WithField("firewallHashes", fw.GetRuleHashes()).
|
||||
WithField("oldFirewallHashes", oldFw.GetRuleHashes()).
|
||||
WithField("rulesVersion", fw.rulesVersion).
|
||||
Warn("firewall rulesVersion has overflowed, resetting conntrack")
|
||||
f.l.Warn("firewall rulesVersion has overflowed, resetting conntrack",
|
||||
"firewallHashes", fw.GetRuleHashes(),
|
||||
"oldFirewallHashes", oldFw.GetRuleHashes(),
|
||||
"rulesVersion", fw.rulesVersion,
|
||||
)
|
||||
} else {
|
||||
fw.Conntrack = conntrack
|
||||
}
|
||||
@@ -360,10 +477,11 @@ func (f *Interface) reloadFirewall(c *config.C) {
|
||||
f.firewall = fw
|
||||
|
||||
oldFw.Destroy()
|
||||
f.l.WithField("firewallHashes", fw.GetRuleHashes()).
|
||||
WithField("oldFirewallHashes", oldFw.GetRuleHashes()).
|
||||
WithField("rulesVersion", fw.rulesVersion).
|
||||
Info("New firewall has been installed")
|
||||
f.l.Info("New firewall has been installed",
|
||||
"firewallHashes", fw.GetRuleHashes(),
|
||||
"oldFirewallHashes", oldFw.GetRuleHashes(),
|
||||
"rulesVersion", fw.rulesVersion,
|
||||
)
|
||||
}
|
||||
|
||||
func (f *Interface) reloadSendRecvError(c *config.C) {
|
||||
@@ -385,8 +503,7 @@ func (f *Interface) reloadSendRecvError(c *config.C) {
|
||||
}
|
||||
}
|
||||
|
||||
f.l.WithField("sendRecvError", f.sendRecvErrorConfig.String()).
|
||||
Info("Loaded send_recv_error config")
|
||||
f.l.Info("Loaded send_recv_error config", "sendRecvError", f.sendRecvErrorConfig.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -409,8 +526,7 @@ func (f *Interface) reloadAcceptRecvError(c *config.C) {
|
||||
}
|
||||
}
|
||||
|
||||
f.l.WithField("acceptRecvError", f.acceptRecvErrorConfig.String()).
|
||||
Info("Loaded accept_recv_error config")
|
||||
f.l.Info("Loaded accept_recv_error config", "acceptRecvError", f.acceptRecvErrorConfig.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -434,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) {
|
||||
ticker := time.NewTicker(i)
|
||||
defer ticker.Stop()
|
||||
@@ -477,23 +607,23 @@ func (f *Interface) GetCertState() *CertState {
|
||||
}
|
||||
|
||||
func (f *Interface) Close() error {
|
||||
var errs []error
|
||||
f.closed.Store(true)
|
||||
|
||||
for _, u := range f.writers {
|
||||
// Release the udp readers
|
||||
for i, u := range f.writers {
|
||||
err := u.Close()
|
||||
if err != nil {
|
||||
f.l.WithError(err).Error("Error while closing udp socket")
|
||||
}
|
||||
}
|
||||
for i, r := range f.readers {
|
||||
if i == 0 {
|
||||
continue // f.readers[0] is f.inside, which we want to save for last
|
||||
}
|
||||
if err := r.Close(); err != nil {
|
||||
f.l.WithError(err).Error("Error while closing tun reader")
|
||||
f.l.Error("Error while closing udp socket", "error", err, "writer", i)
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Release the tun device
|
||||
return f.inside.Close()
|
||||
// Release the tun device (closing the tun also closes all readers)
|
||||
closeErr := f.inside.Close()
|
||||
if closeErr != nil {
|
||||
errs = append(errs, closeErr)
|
||||
}
|
||||
f.wg.Done()
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
+167
-111
@@ -5,6 +5,7 @@ import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
@@ -14,11 +15,10 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/gaissmai/bart"
|
||||
"github.com/rcrowley/go-metrics"
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/cert"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/header"
|
||||
"github.com/slackhq/nebula/logging"
|
||||
"github.com/slackhq/nebula/udp"
|
||||
"github.com/slackhq/nebula/util"
|
||||
)
|
||||
@@ -34,7 +34,6 @@ type LightHouse struct {
|
||||
|
||||
myVpnNetworks []netip.Prefix
|
||||
myVpnNetworksTable *bart.Lite
|
||||
punchConn udp.Conn
|
||||
punchy *Punchy
|
||||
|
||||
// Local cache of answers from light houses
|
||||
@@ -69,18 +68,18 @@ type LightHouse struct {
|
||||
// Addr's of relays that can be used by peers to access me
|
||||
relaysForMe atomic.Pointer[[]netip.Addr]
|
||||
|
||||
queryChan chan netip.Addr
|
||||
updateTrigger chan struct{}
|
||||
queryChan chan netip.Addr
|
||||
|
||||
calculatedRemotes atomic.Pointer[bart.Table[[]*calculatedRemote]] // Maps VpnAddr to []*calculatedRemote
|
||||
|
||||
metrics *MessageMetrics
|
||||
metricHolepunchTx metrics.Counter
|
||||
l *logrus.Logger
|
||||
metrics *MessageMetrics
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
// NewLightHouseFromConfig will build a Lighthouse struct from the values provided in the config object
|
||||
// addrMap should be nil unless this is during a config reload
|
||||
func NewLightHouseFromConfig(ctx context.Context, l *logrus.Logger, c *config.C, cs *CertState, pc udp.Conn, p *Punchy) (*LightHouse, error) {
|
||||
func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, cs *CertState, pc udp.Conn, p *Punchy) (*LightHouse, error) {
|
||||
amLighthouse := c.GetBool("lighthouse.am_lighthouse", false)
|
||||
nebulaPort := uint32(c.GetInt("listen.port", 0))
|
||||
if amLighthouse && nebulaPort == 0 {
|
||||
@@ -103,8 +102,8 @@ func NewLightHouseFromConfig(ctx context.Context, l *logrus.Logger, c *config.C,
|
||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||
addrMap: make(map[netip.Addr]*RemoteList),
|
||||
nebulaPort: nebulaPort,
|
||||
punchConn: pc,
|
||||
punchy: p,
|
||||
updateTrigger: make(chan struct{}, 1),
|
||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||
l: l,
|
||||
}
|
||||
@@ -115,9 +114,6 @@ func NewLightHouseFromConfig(ctx context.Context, l *logrus.Logger, c *config.C,
|
||||
|
||||
if c.GetBool("stats.lighthouse_metrics", false) {
|
||||
h.metrics = newLighthouseMetrics()
|
||||
h.metricHolepunchTx = metrics.GetOrRegisterCounter("messages.tx.holepunch", nil)
|
||||
} else {
|
||||
h.metricHolepunchTx = metrics.NilCounter{}
|
||||
}
|
||||
|
||||
err := h.reload(c, true)
|
||||
@@ -131,7 +127,7 @@ func NewLightHouseFromConfig(ctx context.Context, l *logrus.Logger, c *config.C,
|
||||
case *util.ContextualError:
|
||||
v.Log(l)
|
||||
case error:
|
||||
l.WithError(err).Error("failed to reload lighthouse")
|
||||
l.Error("failed to reload lighthouse", "error", err)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -203,8 +199,10 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
||||
//TODO: we could technically insert all returned addrs instead of just the first one if a dns lookup was used
|
||||
addr := addrs[0].Unmap()
|
||||
if lh.myVpnNetworksTable.Contains(addr) {
|
||||
lh.l.WithField("addr", rawAddr).WithField("entry", i+1).
|
||||
Warn("Ignoring lighthouse.advertise_addrs report because it is within the nebula network range")
|
||||
lh.l.Warn("Ignoring lighthouse.advertise_addrs report because it is within the nebula network range",
|
||||
"addr", rawAddr,
|
||||
"entry", i+1,
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -222,7 +220,9 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
||||
lh.interval.Store(int64(c.GetInt("lighthouse.interval", 10)))
|
||||
|
||||
if !initial {
|
||||
lh.l.Infof("lighthouse.interval changed to %v", lh.interval.Load())
|
||||
lh.l.Info("lighthouse.interval changed",
|
||||
"interval", lh.interval.Load(),
|
||||
)
|
||||
|
||||
if lh.updateCancel != nil {
|
||||
// May not always have a running routine
|
||||
@@ -316,6 +316,7 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
||||
if !initial {
|
||||
//NOTE: we are not tearing down existing lighthouse connections because they might be used for non lighthouse traffic
|
||||
lh.l.Info("lighthouse.hosts has changed")
|
||||
lh.TriggerUpdate()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -333,9 +334,12 @@ func (lh *LightHouse) reload(c *config.C, initial bool) error {
|
||||
for _, v := range c.GetStringSlice("relay.relays", nil) {
|
||||
configRIP, err := netip.ParseAddr(v)
|
||||
if err != nil {
|
||||
lh.l.WithField("relay", v).WithError(err).Warn("Parse relay from config failed")
|
||||
lh.l.Warn("Parse relay from config failed",
|
||||
"relay", v,
|
||||
"error", err,
|
||||
)
|
||||
} else {
|
||||
lh.l.WithField("relay", v).Info("Read relay from config")
|
||||
lh.l.Info("Read relay from config", "relay", v)
|
||||
relaysForMe = append(relaysForMe, configRIP)
|
||||
}
|
||||
}
|
||||
@@ -360,8 +364,10 @@ func (lh *LightHouse) parseLighthouses(c *config.C) ([]netip.Addr, error) {
|
||||
}
|
||||
|
||||
if !lh.myVpnNetworksTable.Contains(addr) {
|
||||
lh.l.WithFields(m{"vpnAddr": addr, "networks": lh.myVpnNetworks}).
|
||||
Warn("lighthouse host is not within our networks, lighthouse functionality will work but layer 3 network traffic to the lighthouse will not")
|
||||
lh.l.Warn("lighthouse host is not within our networks, lighthouse functionality will work but layer 3 network traffic to the lighthouse will not",
|
||||
"vpnAddr", addr,
|
||||
"networks", lh.myVpnNetworks,
|
||||
)
|
||||
}
|
||||
out[i] = addr
|
||||
}
|
||||
@@ -432,8 +438,11 @@ func (lh *LightHouse) loadStaticMap(c *config.C, staticList map[netip.Addr]struc
|
||||
}
|
||||
|
||||
if !lh.myVpnNetworksTable.Contains(vpnAddr) {
|
||||
lh.l.WithFields(m{"vpnAddr": vpnAddr, "networks": lh.myVpnNetworks, "entry": i + 1}).
|
||||
Warn("static_host_map key is not within our networks, layer 3 network traffic to this host will not work")
|
||||
lh.l.Warn("static_host_map key is not within our networks, layer 3 network traffic to this host will not work",
|
||||
"vpnAddr", vpnAddr,
|
||||
"networks", lh.myVpnNetworks,
|
||||
"entry", i+1,
|
||||
)
|
||||
}
|
||||
|
||||
vals, ok := v.([]any)
|
||||
@@ -534,12 +543,13 @@ func (lh *LightHouse) DeleteVpnAddrs(allVpnAddrs []netip.Addr) {
|
||||
lh.Lock()
|
||||
rm, ok := lh.addrMap[allVpnAddrs[0]]
|
||||
if ok {
|
||||
debugEnabled := lh.l.Enabled(context.Background(), slog.LevelDebug)
|
||||
for _, addr := range allVpnAddrs {
|
||||
srm := lh.addrMap[addr]
|
||||
if srm == rm {
|
||||
delete(lh.addrMap, addr)
|
||||
if lh.l.Level >= logrus.DebugLevel {
|
||||
lh.l.Debugf("deleting %s from lighthouse.", addr)
|
||||
if debugEnabled {
|
||||
lh.l.Debug("deleting from lighthouse", "vpnAddr", addr)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -656,9 +666,12 @@ func (lh *LightHouse) unlockedGetRemoteList(allAddrs []netip.Addr) *RemoteList {
|
||||
|
||||
func (lh *LightHouse) shouldAdd(vpnAddrs []netip.Addr, to netip.Addr) bool {
|
||||
allow := lh.GetRemoteAllowList().AllowAll(vpnAddrs, to)
|
||||
if lh.l.Level >= logrus.TraceLevel {
|
||||
lh.l.WithField("vpnAddrs", vpnAddrs).WithField("udpAddr", to).WithField("allow", allow).
|
||||
Trace("remoteAllowList.Allow")
|
||||
if lh.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
lh.l.Log(context.Background(), logging.LevelTrace, "remoteAllowList.Allow",
|
||||
"vpnAddrs", vpnAddrs,
|
||||
"udpAddr", to,
|
||||
"allow", allow,
|
||||
)
|
||||
}
|
||||
if !allow {
|
||||
return false
|
||||
@@ -675,9 +688,12 @@ func (lh *LightHouse) shouldAdd(vpnAddrs []netip.Addr, to netip.Addr) bool {
|
||||
func (lh *LightHouse) unlockedShouldAddV4(vpnAddr netip.Addr, to *V4AddrPort) bool {
|
||||
udpAddr := protoV4AddrPortToNetAddrPort(to)
|
||||
allow := lh.GetRemoteAllowList().Allow(vpnAddr, udpAddr.Addr())
|
||||
if lh.l.Level >= logrus.TraceLevel {
|
||||
lh.l.WithField("vpnAddr", vpnAddr).WithField("udpAddr", udpAddr).WithField("allow", allow).
|
||||
Trace("remoteAllowList.Allow")
|
||||
if lh.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
lh.l.Log(context.Background(), logging.LevelTrace, "remoteAllowList.Allow",
|
||||
"vpnAddr", vpnAddr,
|
||||
"udpAddr", udpAddr,
|
||||
"allow", allow,
|
||||
)
|
||||
}
|
||||
|
||||
if !allow {
|
||||
@@ -695,9 +711,12 @@ func (lh *LightHouse) unlockedShouldAddV4(vpnAddr netip.Addr, to *V4AddrPort) bo
|
||||
func (lh *LightHouse) unlockedShouldAddV6(vpnAddr netip.Addr, to *V6AddrPort) bool {
|
||||
udpAddr := protoV6AddrPortToNetAddrPort(to)
|
||||
allow := lh.GetRemoteAllowList().Allow(vpnAddr, udpAddr.Addr())
|
||||
if lh.l.Level >= logrus.TraceLevel {
|
||||
lh.l.WithField("vpnAddr", vpnAddr).WithField("udpAddr", udpAddr).WithField("allow", allow).
|
||||
Trace("remoteAllowList.Allow")
|
||||
if lh.l.Enabled(context.Background(), logging.LevelTrace) {
|
||||
lh.l.Log(context.Background(), logging.LevelTrace, "remoteAllowList.Allow",
|
||||
"vpnAddr", vpnAddr,
|
||||
"udpAddr", udpAddr,
|
||||
"allow", allow,
|
||||
)
|
||||
}
|
||||
|
||||
if !allow {
|
||||
@@ -772,8 +791,10 @@ func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
||||
|
||||
if v == cert.Version1 {
|
||||
if !addr.Is4() {
|
||||
lh.l.WithField("queryVpnAddr", addr).WithField("lighthouseAddr", lhVpnAddr).
|
||||
Error("Can't query lighthouse for v6 address using a v1 protocol")
|
||||
lh.l.Error("Can't query lighthouse for v6 address using a v1 protocol",
|
||||
"queryVpnAddr", addr,
|
||||
"lighthouseAddr", lhVpnAddr,
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -784,9 +805,11 @@ func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
||||
|
||||
v1Query, err = msg.Marshal()
|
||||
if err != nil {
|
||||
lh.l.WithError(err).WithField("queryVpnAddr", addr).
|
||||
WithField("lighthouseAddr", lhVpnAddr).
|
||||
Error("Failed to marshal lighthouse v1 query payload")
|
||||
lh.l.Error("Failed to marshal lighthouse v1 query payload",
|
||||
"error", err,
|
||||
"queryVpnAddr", addr,
|
||||
"lighthouseAddr", lhVpnAddr,
|
||||
)
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -801,9 +824,11 @@ func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
||||
|
||||
v2Query, err = msg.Marshal()
|
||||
if err != nil {
|
||||
lh.l.WithError(err).WithField("queryVpnAddr", addr).
|
||||
WithField("lighthouseAddr", lhVpnAddr).
|
||||
Error("Failed to marshal lighthouse v2 query payload")
|
||||
lh.l.Error("Failed to marshal lighthouse v2 query payload",
|
||||
"error", err,
|
||||
"queryVpnAddr", addr,
|
||||
"lighthouseAddr", lhVpnAddr,
|
||||
)
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -812,7 +837,11 @@ func (lh *LightHouse) innerQueryServer(addr netip.Addr, nb, out []byte) {
|
||||
queried++
|
||||
|
||||
} else {
|
||||
lh.l.Debugf("Can not query lighthouse for %v using unknown protocol version: %v", addr, v)
|
||||
lh.l.Debug("unsupported protocol version",
|
||||
"op", "query",
|
||||
"queryVpnAddr", addr,
|
||||
"version", v,
|
||||
)
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -841,11 +870,24 @@ func (lh *LightHouse) StartUpdateWorker() {
|
||||
return
|
||||
case <-clockSource.C:
|
||||
continue
|
||||
case <-lh.updateTrigger:
|
||||
continue
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// TriggerUpdate requests an immediate lighthouse update. This is a non-blocking
|
||||
// operation intended to be called after a handshake completes with a lighthouse,
|
||||
// so the lighthouse has our current addresses without waiting for the next
|
||||
// periodic update.
|
||||
func (lh *LightHouse) TriggerUpdate() {
|
||||
select {
|
||||
case lh.updateTrigger <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (lh *LightHouse) SendUpdate() {
|
||||
var v4 []*V4AddrPort
|
||||
var v6 []*V6AddrPort
|
||||
@@ -891,8 +933,9 @@ func (lh *LightHouse) SendUpdate() {
|
||||
if v == cert.Version1 {
|
||||
if v1Update == nil {
|
||||
if !lh.myVpnNetworks[0].Addr().Is4() {
|
||||
lh.l.WithField("lighthouseAddr", lhVpnAddr).
|
||||
Warn("cannot update lighthouse using v1 protocol without an IPv4 address")
|
||||
lh.l.Warn("cannot update lighthouse using v1 protocol without an IPv4 address",
|
||||
"lighthouseAddr", lhVpnAddr,
|
||||
)
|
||||
continue
|
||||
}
|
||||
var relays []uint32
|
||||
@@ -916,8 +959,10 @@ func (lh *LightHouse) SendUpdate() {
|
||||
|
||||
v1Update, err = msg.Marshal()
|
||||
if err != nil {
|
||||
lh.l.WithError(err).WithField("lighthouseAddr", lhVpnAddr).
|
||||
Error("Error while marshaling for lighthouse v1 update")
|
||||
lh.l.Error("Error while marshaling for lighthouse v1 update",
|
||||
"error", err,
|
||||
"lighthouseAddr", lhVpnAddr,
|
||||
)
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -943,8 +988,10 @@ func (lh *LightHouse) SendUpdate() {
|
||||
|
||||
v2Update, err = msg.Marshal()
|
||||
if err != nil {
|
||||
lh.l.WithError(err).WithField("lighthouseAddr", lhVpnAddr).
|
||||
Error("Error while marshaling for lighthouse v2 update")
|
||||
lh.l.Error("Error while marshaling for lighthouse v2 update",
|
||||
"error", err,
|
||||
"lighthouseAddr", lhVpnAddr,
|
||||
)
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -953,7 +1000,10 @@ func (lh *LightHouse) SendUpdate() {
|
||||
updated++
|
||||
|
||||
} else {
|
||||
lh.l.Debugf("Can not update lighthouse using unknown protocol version: %v", v)
|
||||
lh.l.Debug("unsupported protocol version",
|
||||
"op", "update",
|
||||
"version", v,
|
||||
)
|
||||
continue
|
||||
}
|
||||
}
|
||||
@@ -967,7 +1017,7 @@ type LightHouseHandler struct {
|
||||
out []byte
|
||||
pb []byte
|
||||
meta *NebulaMeta
|
||||
l *logrus.Logger
|
||||
l *slog.Logger
|
||||
}
|
||||
|
||||
func (lh *LightHouse) NewRequestHandler() *LightHouseHandler {
|
||||
@@ -1016,14 +1066,19 @@ func (lhh *LightHouseHandler) HandleRequest(rAddr netip.AddrPort, fromVpnAddrs [
|
||||
n := lhh.resetMeta()
|
||||
err := n.Unmarshal(p)
|
||||
if err != nil {
|
||||
lhh.l.WithError(err).WithField("vpnAddrs", fromVpnAddrs).WithField("udpAddr", rAddr).
|
||||
Error("Failed to unmarshal lighthouse packet")
|
||||
lhh.l.Error("Failed to unmarshal lighthouse packet",
|
||||
"error", err,
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
"udpAddr", rAddr,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
if n.Details == nil {
|
||||
lhh.l.WithField("vpnAddrs", fromVpnAddrs).WithField("udpAddr", rAddr).
|
||||
Error("Invalid lighthouse update")
|
||||
lhh.l.Error("Invalid lighthouse update",
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
"udpAddr", rAddr,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1051,25 +1106,29 @@ func (lhh *LightHouseHandler) HandleRequest(rAddr netip.AddrPort, fromVpnAddrs [
|
||||
func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []netip.Addr, addr netip.AddrPort, w EncWriter) {
|
||||
// Exit if we don't answer queries
|
||||
if !lhh.lh.amLighthouse {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.Debugln("I don't answer queries, but received from: ", addr)
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("I don't answer queries, but received one", "from", addr)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
queryVpnAddr, useVersion, err := n.Details.GetVpnAddrAndVersion()
|
||||
if err != nil {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithField("from", fromVpnAddrs).WithField("details", n.Details).
|
||||
Debugln("Dropping malformed HostQuery")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("Dropping malformed HostQuery",
|
||||
"from", fromVpnAddrs,
|
||||
"details", n.Details,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
if useVersion == cert.Version1 && queryVpnAddr.Is6() {
|
||||
// this case really shouldn't be possible to represent, but reject it anyway.
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithField("vpnAddrs", fromVpnAddrs).WithField("queryVpnAddr", queryVpnAddr).
|
||||
Debugln("invalid vpn addr for v1 handleHostQuery")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("invalid vpn addr for v1 handleHostQuery",
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
"queryVpnAddr", queryVpnAddr,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -1094,7 +1153,10 @@ func (lhh *LightHouseHandler) handleHostQuery(n *NebulaMeta, fromVpnAddrs []neti
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
lhh.l.WithError(err).WithField("vpnAddrs", fromVpnAddrs).Error("Failed to marshal lighthouse host query reply")
|
||||
lhh.l.Error("Failed to marshal lighthouse host query reply",
|
||||
"error", err,
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1122,8 +1184,10 @@ func (lhh *LightHouseHandler) sendHostPunchNotification(n *NebulaMeta, fromVpnAd
|
||||
if ok {
|
||||
whereToPunch = newDest
|
||||
} else {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithField("to", crt.Networks()).Debugln("unable to punch to host, no addresses in common")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("unable to punch to host, no addresses in common",
|
||||
"to", crt.Networks(),
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1149,7 +1213,10 @@ func (lhh *LightHouseHandler) sendHostPunchNotification(n *NebulaMeta, fromVpnAd
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
lhh.l.WithError(err).WithField("vpnAddrs", fromVpnAddrs).Error("Failed to marshal lighthouse host was queried for")
|
||||
lhh.l.Error("Failed to marshal lighthouse host was queried for",
|
||||
"error", err,
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1191,8 +1258,11 @@ func (lhh *LightHouseHandler) coalesceAnswers(v cert.Version, c *cache, n *Nebul
|
||||
n.Details.RelayVpnAddrs = append(n.Details.RelayVpnAddrs, netAddrToProtoAddr(r))
|
||||
}
|
||||
} else {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithField("version", v).Debug("unsupported protocol version")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("unsupported protocol version",
|
||||
"op", "coalesceAnswers",
|
||||
"version", v,
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1205,8 +1275,11 @@ func (lhh *LightHouseHandler) handleHostQueryReply(n *NebulaMeta, fromVpnAddrs [
|
||||
|
||||
certVpnAddr, _, err := n.Details.GetVpnAddrAndVersion()
|
||||
if err != nil {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithError(err).WithField("vpnAddrs", fromVpnAddrs).Error("dropping malformed HostQueryReply")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Error("dropping malformed HostQueryReply",
|
||||
"error", err,
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -1231,8 +1304,8 @@ func (lhh *LightHouseHandler) handleHostQueryReply(n *NebulaMeta, fromVpnAddrs [
|
||||
|
||||
func (lhh *LightHouseHandler) handleHostUpdateNotification(n *NebulaMeta, fromVpnAddrs []netip.Addr, w EncWriter) {
|
||||
if !lhh.lh.amLighthouse {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.Debugln("I am not a lighthouse, do not take host updates: ", fromVpnAddrs)
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("I am not a lighthouse, do not take host updates", "from", fromVpnAddrs)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -1255,8 +1328,11 @@ func (lhh *LightHouseHandler) handleHostUpdateNotification(n *NebulaMeta, fromVp
|
||||
|
||||
//Simple check that the host sent this not someone else, if detailsVpnAddr is filled
|
||||
if detailsVpnAddr.IsValid() && !slices.Contains(fromVpnAddrs, detailsVpnAddr) {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithField("vpnAddrs", fromVpnAddrs).WithField("answer", detailsVpnAddr).Debugln("Host sent invalid update")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("Host sent invalid update",
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
"answer", detailsVpnAddr,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -1278,7 +1354,9 @@ func (lhh *LightHouseHandler) handleHostUpdateNotification(n *NebulaMeta, fromVp
|
||||
switch useVersion {
|
||||
case cert.Version1:
|
||||
if !fromVpnAddrs[0].Is4() {
|
||||
lhh.l.WithField("vpnAddrs", fromVpnAddrs).Error("Can not send HostUpdateNotificationAck for a ipv6 vpn ip in a v1 message")
|
||||
lhh.l.Error("Can not send HostUpdateNotificationAck for a ipv6 vpn ip in a v1 message",
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
)
|
||||
return
|
||||
}
|
||||
vpnAddrB := fromVpnAddrs[0].As4()
|
||||
@@ -1286,13 +1364,16 @@ func (lhh *LightHouseHandler) handleHostUpdateNotification(n *NebulaMeta, fromVp
|
||||
case cert.Version2:
|
||||
// do nothing, we want to send a blank message
|
||||
default:
|
||||
lhh.l.WithField("useVersion", useVersion).Error("invalid protocol version")
|
||||
lhh.l.Error("invalid protocol version", "useVersion", useVersion)
|
||||
return
|
||||
}
|
||||
|
||||
ln, err := n.MarshalTo(lhh.pb)
|
||||
if err != nil {
|
||||
lhh.l.WithError(err).WithField("vpnAddrs", fromVpnAddrs).Error("Failed to marshal lighthouse host update ack")
|
||||
lhh.l.Error("Failed to marshal lighthouse host update ack",
|
||||
"error", err,
|
||||
"vpnAddrs", fromVpnAddrs,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1309,59 +1390,34 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn
|
||||
|
||||
detailsVpnAddr, _, err := n.Details.GetVpnAddrAndVersion()
|
||||
if err != nil {
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.WithField("details", n.Details).WithError(err).Debugln("dropping invalid HostPunchNotification")
|
||||
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
lhh.l.Debug("dropping invalid HostPunchNotification",
|
||||
"details", n.Details,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
empty := []byte{0}
|
||||
punch := func(vpnPeer netip.AddrPort, logVpnAddr netip.Addr) {
|
||||
if !vpnPeer.IsValid() {
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
time.Sleep(lhh.lh.punchy.GetDelay())
|
||||
lhh.lh.metricHolepunchTx.Inc(1)
|
||||
lhh.lh.punchConn.WriteTo(empty, vpnPeer)
|
||||
}()
|
||||
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.Debugf("Punching on %v for %v", vpnPeer, logVpnAddr)
|
||||
}
|
||||
}
|
||||
|
||||
remoteAllowList := lhh.lh.GetRemoteAllowList()
|
||||
for _, a := range n.Details.V4AddrPorts {
|
||||
b := protoV4AddrPortToNetAddrPort(a)
|
||||
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
||||
punch(b, detailsVpnAddr)
|
||||
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
||||
}
|
||||
}
|
||||
|
||||
for _, a := range n.Details.V6AddrPorts {
|
||||
b := protoV6AddrPortToNetAddrPort(a)
|
||||
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
||||
punch(b, detailsVpnAddr)
|
||||
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
||||
}
|
||||
}
|
||||
|
||||
// This sends a nebula test packet to the host trying to contact us. In the case
|
||||
// of a double nat or other difficult scenario, this may help establish
|
||||
// a tunnel.
|
||||
if lhh.lh.punchy.GetRespond() {
|
||||
go func() {
|
||||
time.Sleep(lhh.lh.punchy.GetRespondDelay())
|
||||
if lhh.l.Level >= logrus.DebugLevel {
|
||||
lhh.l.Debugf("Sending a nebula test packet to vpn addr %s", detailsVpnAddr)
|
||||
}
|
||||
//NOTE: we have to allocate a new output buffer here since we are spawning a new goroutine
|
||||
// for each punchBack packet. We should move this into a timerwheel or a single goroutine
|
||||
// managed by a channel.
|
||||
w.SendMessageToVpnAddr(header.Test, header.TestRequest, detailsVpnAddr, []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
||||
}()
|
||||
}
|
||||
// a tunnel. ScheduleRespond is a no-op when punchy.respond is disabled.
|
||||
lhh.lh.punchy.ScheduleRespond(detailsVpnAddr)
|
||||
}
|
||||
|
||||
func protoAddrToNetAddr(addr *Addr) netip.Addr {
|
||||
|
||||
@@ -1,45 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/config"
|
||||
)
|
||||
|
||||
func configLogger(l *logrus.Logger, c *config.C) error {
|
||||
// set up our logging level
|
||||
logLevel, err := logrus.ParseLevel(strings.ToLower(c.GetString("logging.level", "info")))
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s; possible levels: %s", err, logrus.AllLevels)
|
||||
}
|
||||
l.SetLevel(logLevel)
|
||||
|
||||
disableTimestamp := c.GetBool("logging.disable_timestamp", false)
|
||||
timestampFormat := c.GetString("logging.timestamp_format", "")
|
||||
fullTimestamp := (timestampFormat != "")
|
||||
if timestampFormat == "" {
|
||||
timestampFormat = time.RFC3339
|
||||
}
|
||||
|
||||
logFormat := strings.ToLower(c.GetString("logging.format", "text"))
|
||||
switch logFormat {
|
||||
case "text":
|
||||
l.Formatter = &logrus.TextFormatter{
|
||||
TimestampFormat: timestampFormat,
|
||||
FullTimestamp: fullTimestamp,
|
||||
DisableTimestamp: disableTimestamp,
|
||||
}
|
||||
case "json":
|
||||
l.Formatter = &logrus.JSONFormatter{
|
||||
TimestampFormat: timestampFormat,
|
||||
DisableTimestamp: disableTimestamp,
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("unknown log format `%s`. possible formats: %s", logFormat, []string{"text", "json"})
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,233 @@
|
||||
// Package logging wires the nebula runtime-reconfigurable slog handler used
|
||||
// by nebula.Main and the nebula CLI binaries. Callers build a logger with
|
||||
// NewLogger, then call ApplyConfig at startup and from a config reload
|
||||
// callback to push logging.level, logging.format, and
|
||||
// logging.disable_timestamp changes onto the logger without rebuilding it.
|
||||
package logging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config is the subset of *config.C that ApplyConfig reads. Declaring it
|
||||
// here keeps the logging package from depending on config directly, which
|
||||
// would cycle through the shared test helpers (test.NewLogger imports
|
||||
// logging, and config's tests import test). *config.C satisfies this
|
||||
// interface structurally with no adapter.
|
||||
type Config interface {
|
||||
GetString(key, def string) string
|
||||
GetBool(key string, def bool) bool
|
||||
}
|
||||
|
||||
// LevelTrace is a custom slog level below Debug, used when logging.level is
|
||||
// "trace". slog has no builtin trace level; the value is one step below
|
||||
// slog.LevelDebug in slog's 4-point spacing.
|
||||
const LevelTrace = slog.Level(-8)
|
||||
|
||||
// NewLogger returns a *slog.Logger whose level, format, and timestamp
|
||||
// emission can be reconfigured at runtime via ApplyConfig and the SSH debug
|
||||
// commands. The default configuration is info-level text output so log
|
||||
// calls made before ApplyConfig runs still produce output. Timestamps
|
||||
// follow slog's default RFC3339Nano format; set logging.disable_timestamp
|
||||
// in config to suppress them.
|
||||
//
|
||||
// ApplyConfig and the SSH commands discover the reconfig surface via
|
||||
// structural type-assertion on l.Handler(), so replacement implementations
|
||||
// (tests, platform-specific sinks) need only implement the subset of
|
||||
// {SetLevel(slog.Level), SetFormat(string) error, SetDisableTimestamp(bool)}
|
||||
// they care about. Callers that pass a plain *slog.Logger without these
|
||||
// methods get a silent no-op; reconfiguration is always opt-in.
|
||||
func NewLogger(w io.Writer) *slog.Logger {
|
||||
return slog.New(NewHandler(w))
|
||||
}
|
||||
|
||||
// NewHandler builds the *Handler that NewLogger wraps. Exported for
|
||||
// platform-specific sinks (notably cmd/nebula-service/logs_windows.go)
|
||||
// that want to wrap the handler with extra behavior, such as tagging each
|
||||
// record with its Event Log severity, while still benefiting from all the
|
||||
// level / format / timestamp / WithAttrs machinery implemented here.
|
||||
func NewHandler(w io.Writer) *Handler {
|
||||
root := &handlerRoot{}
|
||||
root.level.Set(slog.LevelInfo)
|
||||
opts := &slog.HandlerOptions{Level: &root.level}
|
||||
return &Handler{
|
||||
root: root,
|
||||
text: slog.NewTextHandler(w, opts),
|
||||
json: slog.NewJSONHandler(w, opts),
|
||||
}
|
||||
}
|
||||
|
||||
// handlerRoot carries the reconfiguration state shared by every logger
|
||||
// derived from a NewHandler call. All fields are consulted on the log
|
||||
// path and updated lock-free.
|
||||
type handlerRoot struct {
|
||||
level slog.LevelVar
|
||||
disableTimestamp atomic.Bool
|
||||
// jsonMode picks which of the pre-derived inner handlers Handler.Handle
|
||||
// dispatches to. Flipping it propagates instantly to every derived logger
|
||||
// without rebuilding or chain-replaying anything.
|
||||
jsonMode atomic.Bool
|
||||
}
|
||||
|
||||
// Handler is the slog.Handler returned by NewHandler. It holds two
|
||||
// pre-derived slog handlers -- one text, one json -- both built from the
|
||||
// same accumulated WithAttrs/WithGroup state. Handle picks which one to
|
||||
// dispatch to based on handlerRoot.jsonMode, so a SetFormat call takes
|
||||
// effect immediately across the whole process without having to rebuild
|
||||
// any derived loggers.
|
||||
type Handler struct {
|
||||
root *handlerRoot
|
||||
text slog.Handler
|
||||
json slog.Handler
|
||||
}
|
||||
|
||||
func (h *Handler) Enabled(_ context.Context, l slog.Level) bool {
|
||||
return h.root.level.Level() <= l
|
||||
}
|
||||
|
||||
func (h *Handler) Handle(ctx context.Context, r slog.Record) error {
|
||||
if h.root.disableTimestamp.Load() {
|
||||
r.Time = time.Time{}
|
||||
}
|
||||
if h.root.jsonMode.Load() {
|
||||
return h.json.Handle(ctx, r)
|
||||
}
|
||||
return h.text.Handle(ctx, r)
|
||||
}
|
||||
|
||||
func (h *Handler) WithAttrs(attrs []slog.Attr) slog.Handler {
|
||||
if len(attrs) == 0 {
|
||||
return h
|
||||
}
|
||||
return &Handler{
|
||||
root: h.root,
|
||||
text: h.text.WithAttrs(attrs),
|
||||
json: h.json.WithAttrs(attrs),
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) WithGroup(name string) slog.Handler {
|
||||
if name == "" {
|
||||
return h
|
||||
}
|
||||
return &Handler{
|
||||
root: h.root,
|
||||
text: h.text.WithGroup(name),
|
||||
json: h.json.WithGroup(name),
|
||||
}
|
||||
}
|
||||
|
||||
// SetLevel updates the effective log level. Propagates to every derived
|
||||
// logger via the shared LevelVar.
|
||||
func (h *Handler) SetLevel(level slog.Level) { h.root.level.Set(level) }
|
||||
|
||||
// GetLevel reports the current log level.
|
||||
func (h *Handler) GetLevel() slog.Level { return h.root.level.Level() }
|
||||
|
||||
// SetFormat flips the output format atomically. Valid formats are "text"
|
||||
// and "json". Every derived logger sees the new format on its next Handle
|
||||
// call; no rebuild or registration is required.
|
||||
func (h *Handler) SetFormat(format string) error {
|
||||
switch format {
|
||||
case "text":
|
||||
h.root.jsonMode.Store(false)
|
||||
case "json":
|
||||
h.root.jsonMode.Store(true)
|
||||
default:
|
||||
return fmt.Errorf("unknown log format `%s`. possible formats: %s", format, []string{"text", "json"})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetFormat reports the currently selected format name.
|
||||
func (h *Handler) GetFormat() string {
|
||||
if h.root.jsonMode.Load() {
|
||||
return "json"
|
||||
}
|
||||
return "text"
|
||||
}
|
||||
|
||||
// SetDisableTimestamp toggles whether Handle zeroes r.Time before
|
||||
// dispatching (slog's builtin text/json handlers skip emitting the time
|
||||
// attribute on a zero time).
|
||||
func (h *Handler) SetDisableTimestamp(v bool) { h.root.disableTimestamp.Store(v) }
|
||||
|
||||
// ApplyConfig reads logging.level, logging.format, and (optionally)
|
||||
// logging.disable_timestamp from c and applies them to l. The reconfig
|
||||
// surface is discovered via structural type-assertion on l.Handler(), so
|
||||
// foreign handlers silently opt out of whichever capabilities they do not
|
||||
// implement.
|
||||
//
|
||||
// nebula.Main does NOT call this function on your behalf; callers that want
|
||||
// config-driven log level / format / timestamp updates invoke it at
|
||||
// startup and register it as a reload callback themselves. This keeps the
|
||||
// library from mutating an embedder's logger without their say-so.
|
||||
func ApplyConfig(l *slog.Logger, c Config) error {
|
||||
h := l.Handler()
|
||||
|
||||
lvl, err := ParseLevel(strings.ToLower(c.GetString("logging.level", "info")))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if ls, ok := h.(interface{ SetLevel(slog.Level) }); ok {
|
||||
ls.SetLevel(lvl)
|
||||
}
|
||||
|
||||
format := strings.ToLower(c.GetString("logging.format", "text"))
|
||||
if fs, ok := h.(interface{ SetFormat(string) error }); ok {
|
||||
if err := fs.SetFormat(format); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if ts, ok := h.(interface{ SetDisableTimestamp(bool) }); ok {
|
||||
ts.SetDisableTimestamp(c.GetBool("logging.disable_timestamp", false))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ParseLevel converts a config-string level name ("trace", "debug", "info",
|
||||
// "warn"/"warning", "error", "fatal"/"panic") to a slog.Level. "fatal" and
|
||||
// "panic" are accepted for backwards compatibility with pre-slog configs
|
||||
// and both map to slog.LevelError.
|
||||
func ParseLevel(s string) (slog.Level, error) {
|
||||
switch s {
|
||||
case "trace":
|
||||
return LevelTrace, nil
|
||||
case "debug":
|
||||
return slog.LevelDebug, nil
|
||||
case "info":
|
||||
return slog.LevelInfo, nil
|
||||
case "warn", "warning":
|
||||
return slog.LevelWarn, nil
|
||||
case "error":
|
||||
return slog.LevelError, nil
|
||||
case "fatal", "panic":
|
||||
return slog.LevelError, nil
|
||||
default:
|
||||
return 0, fmt.Errorf("not a valid logging level: %q", s)
|
||||
}
|
||||
}
|
||||
|
||||
// LevelName returns a human-readable name for a slog.Level matching the
|
||||
// strings accepted by ParseLevel.
|
||||
func LevelName(l slog.Level) string {
|
||||
switch {
|
||||
case l <= LevelTrace:
|
||||
return "trace"
|
||||
case l <= slog.LevelDebug:
|
||||
return "debug"
|
||||
case l <= slog.LevelInfo:
|
||||
return "info"
|
||||
case l <= slog.LevelWarn:
|
||||
return "warn"
|
||||
default:
|
||||
return "error"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package logging
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"log/slog"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// BenchmarkLogger_* compare the handler returned by NewLogger against a
|
||||
// stock slog text handler. The key thing we care about is the per-log
|
||||
// cost on a logger that has been derived via .With(), because that is the
|
||||
// shape subsystems store on their structs (HostInfo.logger(),
|
||||
// lh.l.With("subsystem", ...), etc.) and call from hot paths.
|
||||
|
||||
func BenchmarkLogger_Stock_RootInfo(b *testing.B) {
|
||||
l := slog.New(slog.DiscardHandler)
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
l.Info("hello", "i", i)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLogger_Nebula_RootInfo(b *testing.B) {
|
||||
l := NewLogger(io.Discard)
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
l.Info("hello", "i", i)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLogger_Stock_DerivedInfo(b *testing.B) {
|
||||
l := slog.New(slog.DiscardHandler).With(
|
||||
"subsystem", "bench",
|
||||
"localIndex", 1234,
|
||||
)
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
l.Info("hello", "i", i)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLogger_Nebula_DerivedInfo(b *testing.B) {
|
||||
l := NewLogger(io.Discard).With(
|
||||
"subsystem", "bench",
|
||||
"localIndex", 1234,
|
||||
)
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
l.Info("hello", "i", i)
|
||||
}
|
||||
}
|
||||
|
||||
// Gated-off-path benchmarks: mimic the typical hot-path shape
|
||||
// `if l.Enabled(ctx, slog.LevelDebug) { ... }` where the log is gated below
|
||||
// the active level. This is the dominant pattern in inside.go/outside.go and
|
||||
// what we pay on every packet.
|
||||
func BenchmarkLogger_Stock_DerivedEnabledGateMiss(b *testing.B) {
|
||||
l := slog.New(slog.DiscardHandler).With(
|
||||
"subsystem", "bench",
|
||||
"localIndex", 1234,
|
||||
)
|
||||
ctx := context.Background()
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if l.Enabled(ctx, slog.LevelDebug) {
|
||||
l.Debug("hello", "i", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkLogger_Nebula_DerivedEnabledGateMiss(b *testing.B) {
|
||||
l := NewLogger(io.Discard).With(
|
||||
"subsystem", "bench",
|
||||
"localIndex", 1234,
|
||||
)
|
||||
ctx := context.Background()
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
if l.Enabled(ctx, slog.LevelDebug) {
|
||||
l.Debug("hello", "i", i)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -3,13 +3,16 @@ package nebula
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
_ "net/http/pprof"
|
||||
"net/netip"
|
||||
"runtime"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/config"
|
||||
"github.com/slackhq/nebula/overlay"
|
||||
"github.com/slackhq/nebula/sshd"
|
||||
@@ -20,7 +23,7 @@ import (
|
||||
|
||||
type m = map[string]any
|
||||
|
||||
func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logger, deviceFactory overlay.DeviceFactory) (retcon *Control, reterr error) {
|
||||
func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, deviceFactory overlay.DeviceFactory) (retcon *Control, reterr error) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Automatically cancel the context if Main returns an error, to signal all created goroutines to quit.
|
||||
defer func() {
|
||||
@@ -33,10 +36,8 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
buildVersion = moduleVersion()
|
||||
}
|
||||
|
||||
l := logger
|
||||
l.Formatter = &logrus.TextFormatter{
|
||||
FullTimestamp: true,
|
||||
}
|
||||
//todo no merge
|
||||
go http.ListenAndServe(":6060", nil)
|
||||
|
||||
// Print the config if in test, the exit comes later
|
||||
if configTest {
|
||||
@@ -46,21 +47,9 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
}
|
||||
|
||||
// Print the final config
|
||||
l.Println(string(b))
|
||||
l.Info(string(b))
|
||||
}
|
||||
|
||||
err := configLogger(l, c)
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Failed to configure the logger", err)
|
||||
}
|
||||
|
||||
c.RegisterReloadCallback(func(c *config.C) {
|
||||
err := configLogger(l, c)
|
||||
if err != nil {
|
||||
l.WithError(err).Error("Failed to configure the logger")
|
||||
}
|
||||
})
|
||||
|
||||
pki, err := NewPKIFromConfig(l, c)
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Failed to load PKI from config", err)
|
||||
@@ -70,9 +59,9 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Error while loading firewall rules", err)
|
||||
}
|
||||
l.WithField("firewallHashes", fw.GetRuleHashes()).Info("Firewall started")
|
||||
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
||||
|
||||
ssh, err := sshd.NewSSHServer(l.WithField("subsystem", "sshd"))
|
||||
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
||||
}
|
||||
@@ -81,7 +70,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
if c.GetBool("sshd.enabled", false) {
|
||||
sshStart, err = configSSH(l, ssh, c)
|
||||
if err != nil {
|
||||
l.WithError(err).Warn("Failed to configure sshd, ssh debugging will not be available")
|
||||
l.Warn("Failed to configure sshd, ssh debugging will not be available", "error", err)
|
||||
sshStart = nil
|
||||
}
|
||||
}
|
||||
@@ -99,7 +88,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
routines = 1
|
||||
}
|
||||
if routines > 1 {
|
||||
l.WithField("routines", routines).Info("Using multiple routines")
|
||||
l.Info("Using multiple routines", "routines", routines)
|
||||
}
|
||||
} else {
|
||||
// deprecated and undocumented
|
||||
@@ -107,7 +96,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
udpQueues := c.GetInt("listen.routines", 1)
|
||||
routines = max(tunQueues, udpQueues)
|
||||
if routines != 1 {
|
||||
l.WithField("routines", routines).Warn("Setting tun.routines and listen.routines is deprecated. Use `routines` instead")
|
||||
l.Warn("Setting tun.routines and listen.routines is deprecated. Use `routines` instead", "routines", routines)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -120,7 +109,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
conntrackCacheTimeout = 1 * time.Second
|
||||
}
|
||||
if conntrackCacheTimeout > 0 {
|
||||
l.WithField("duration", conntrackCacheTimeout).Info("Using routine-local conntrack cache")
|
||||
l.Info("Using routine-local conntrack cache", "duration", conntrackCacheTimeout)
|
||||
}
|
||||
|
||||
var tun overlay.Device
|
||||
@@ -166,7 +155,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
}
|
||||
|
||||
for i := 0; i < routines; i++ {
|
||||
l.Infof("listening on %v", netip.AddrPortFrom(listenHost, uint16(port)))
|
||||
l.Info("listening", "addr", netip.AddrPortFrom(listenHost, uint16(port)))
|
||||
udpServer, err := udp.NewListener(l, listenHost, port, routines > 1, c.GetInt("listen.batch", 64))
|
||||
if err != nil {
|
||||
return nil, util.NewContextualError("Failed to open udp listener", m{"queue": i}, err)
|
||||
@@ -187,7 +176,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
}
|
||||
|
||||
hostMap := NewHostMapFromConfig(l, c)
|
||||
punchy := NewPunchyFromConfig(l, c)
|
||||
punchy := NewPunchyFromConfig(l, c, udpConns[0])
|
||||
connManager := newConnectionManagerFromConfig(l, c, hostMap, punchy)
|
||||
lightHouse, err := NewLightHouseFromConfig(ctx, l, c, pki.getCertState(), udpConns[0], punchy)
|
||||
if err != nil {
|
||||
@@ -201,27 +190,19 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
messageMetrics = newMessageMetricsOnlyRecvError()
|
||||
}
|
||||
|
||||
useRelays := c.GetBool("relay.use_relays", DefaultUseRelays) && !c.GetBool("relay.am_relay", false)
|
||||
|
||||
handshakeConfig := HandshakeConfig{
|
||||
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
||||
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
||||
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
||||
useRelays: useRelays,
|
||||
|
||||
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
||||
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
||||
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
||||
messageMetrics: messageMetrics,
|
||||
}
|
||||
|
||||
handshakeManager := NewHandshakeManager(l, hostMap, lightHouse, udpConns[0], handshakeConfig)
|
||||
lightHouse.handshakeTrigger = handshakeManager.trigger
|
||||
|
||||
serveDns := false
|
||||
if c.GetBool("lighthouse.serve_dns", false) {
|
||||
if c.GetBool("lighthouse.am_lighthouse", false) {
|
||||
serveDns = true
|
||||
} else {
|
||||
l.Warn("DNS server refusing to run because this host is not a lighthouse.")
|
||||
}
|
||||
ds, err := newDnsServerFromConfig(ctx, l, pki.getCertState(), hostMap, c)
|
||||
if err != nil {
|
||||
l.Warn("Failed to start DNS responder", "error", err)
|
||||
}
|
||||
|
||||
ifConfig := &InterfaceConfig{
|
||||
@@ -230,7 +211,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
Outside: udpConns[0],
|
||||
pki: pki,
|
||||
Firewall: fw,
|
||||
ServeDns: serveDns,
|
||||
DnsServer: ds,
|
||||
HandshakeManager: handshakeManager,
|
||||
connectionManager: connManager,
|
||||
lightHouse: lightHouse,
|
||||
@@ -245,6 +226,7 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||
punchy: punchy,
|
||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||
CpuAffinity: parseCpuAffinity(c, l, routines),
|
||||
l: l,
|
||||
}
|
||||
|
||||
@@ -262,12 +244,15 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
ifce.reloadDisconnectInvalid(c)
|
||||
ifce.reloadSendRecvError(c)
|
||||
ifce.reloadAcceptRecvError(c)
|
||||
ifce.reloadEcn(c)
|
||||
|
||||
handshakeManager.f = ifce
|
||||
go handshakeManager.Run(ctx)
|
||||
|
||||
punchy.Start(ctx, ifce, hostMap, lightHouse)
|
||||
}
|
||||
|
||||
statsStart, err := startStats(l, c, buildVersion, configTest)
|
||||
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
||||
if err != nil {
|
||||
return nil, util.ContextualizeIfNeeded("Failed to start stats emitter", err)
|
||||
}
|
||||
@@ -280,26 +265,67 @@ func Main(c *config.C, configTest bool, buildVersion string, logger *logrus.Logg
|
||||
|
||||
attachCommands(l, c, ssh, ifce)
|
||||
|
||||
// Start DNS server last to allow using the nebula IP as lighthouse.dns.host
|
||||
var dnsStart func()
|
||||
if lightHouse.amLighthouse && serveDns {
|
||||
l.Debugln("Starting dns server")
|
||||
dnsStart = dnsMain(l, pki.getCertState(), hostMap, c)
|
||||
}
|
||||
|
||||
return &Control{
|
||||
ifce,
|
||||
l,
|
||||
ctx,
|
||||
cancel,
|
||||
sshStart,
|
||||
statsStart,
|
||||
dnsStart,
|
||||
lightHouse.StartUpdateWorker,
|
||||
connManager.Start,
|
||||
state: StateReady,
|
||||
f: ifce,
|
||||
l: l,
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
sshStart: sshStart,
|
||||
statsStart: stats.Start,
|
||||
dnsStart: ds.Start,
|
||||
lighthouseStart: lightHouse.StartUpdateWorker,
|
||||
connectionManagerStart: connManager.Start,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// parseCpuAffinity reads `tun.cpu_affinity` from the config — a list of
|
||||
// integer CPU IDs, one per TUN reader goroutine. Empty / unset returns nil
|
||||
// (listenIn falls back to its default `i % NumCPU` pinning). Length
|
||||
// mismatch with `routines` is a warning, not an error: shorter lists are
|
||||
// modulo-cycled across queues, longer lists' tail is ignored. Invalid
|
||||
// entries (non-integer, out of range) are also a warning and disable the
|
||||
// override entirely so we don't silently pin to the wrong CPU.
|
||||
func parseCpuAffinity(c *config.C, l *slog.Logger, routines int) []int {
|
||||
raw := c.Get("tun.cpu_affinity")
|
||||
if raw == nil {
|
||||
return nil
|
||||
}
|
||||
rv, ok := raw.([]any)
|
||||
if !ok {
|
||||
l.Warn("tun.cpu_affinity must be a list of integers; ignoring", "value", raw)
|
||||
return nil
|
||||
}
|
||||
nCPU := runtime.NumCPU()
|
||||
cpus := make([]int, 0, len(rv))
|
||||
for i, e := range rv {
|
||||
var cpu int
|
||||
switch v := e.(type) {
|
||||
case int:
|
||||
cpu = v
|
||||
case int64:
|
||||
cpu = int(v)
|
||||
case float64:
|
||||
cpu = int(v)
|
||||
default:
|
||||
l.Warn("tun.cpu_affinity entry not an integer; ignoring affinity",
|
||||
"index", i, "value", e)
|
||||
return nil
|
||||
}
|
||||
if cpu < 0 || cpu >= nCPU {
|
||||
l.Warn("tun.cpu_affinity entry out of range; ignoring affinity",
|
||||
"index", i, "cpu", cpu, "num_cpu", nCPU)
|
||||
return nil
|
||||
}
|
||||
cpus = append(cpus, cpu)
|
||||
}
|
||||
if len(cpus) != routines {
|
||||
l.Warn("tun.cpu_affinity length doesn't match routines; queues will modulo-cycle through the list",
|
||||
"affinity_len", len(cpus), "routines", routines)
|
||||
}
|
||||
return cpus
|
||||
}
|
||||
|
||||
func moduleVersion() string {
|
||||
info, ok := debug.ReadBuildInfo()
|
||||
if !ok {
|
||||
|
||||
@@ -13,6 +13,8 @@ type MessageMetrics struct {
|
||||
|
||||
rxUnknown metrics.Counter
|
||||
txUnknown metrics.Counter
|
||||
|
||||
rxInvalid metrics.Counter
|
||||
}
|
||||
|
||||
func (m *MessageMetrics) Rx(t header.MessageType, s header.MessageSubType, i int64) {
|
||||
@@ -33,6 +35,11 @@ func (m *MessageMetrics) Tx(t header.MessageType, s header.MessageSubType, i int
|
||||
}
|
||||
}
|
||||
}
|
||||
func (m *MessageMetrics) RxInvalid(i int64) {
|
||||
if m != nil && m.rxInvalid != nil {
|
||||
m.rxInvalid.Inc(i)
|
||||
}
|
||||
}
|
||||
|
||||
func newMessageMetrics() *MessageMetrics {
|
||||
gen := func(t string) [][]metrics.Counter {
|
||||
@@ -56,6 +63,7 @@ func newMessageMetrics() *MessageMetrics {
|
||||
|
||||
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
||||
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
||||
rxInvalid: metrics.GetOrRegisterCounter("messages.rx.invalid", nil),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+45
-632
@@ -124,7 +124,7 @@ func (x NebulaControl_MessageType) String() string {
|
||||
}
|
||||
|
||||
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
||||
return fileDescriptor_2d65afa7693df5ef, []int{8, 0}
|
||||
return fileDescriptor_2d65afa7693df5ef, []int{6, 0}
|
||||
}
|
||||
|
||||
type NebulaMeta struct {
|
||||
@@ -489,142 +489,6 @@ func (m *NebulaPing) GetTime() uint64 {
|
||||
return 0
|
||||
}
|
||||
|
||||
type NebulaHandshake struct {
|
||||
Details *NebulaHandshakeDetails `protobuf:"bytes,1,opt,name=Details,proto3" json:"Details,omitempty"`
|
||||
Hmac []byte `protobuf:"bytes,2,opt,name=Hmac,proto3" json:"Hmac,omitempty"`
|
||||
}
|
||||
|
||||
func (m *NebulaHandshake) Reset() { *m = NebulaHandshake{} }
|
||||
func (m *NebulaHandshake) String() string { return proto.CompactTextString(m) }
|
||||
func (*NebulaHandshake) ProtoMessage() {}
|
||||
func (*NebulaHandshake) Descriptor() ([]byte, []int) {
|
||||
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
||||
}
|
||||
func (m *NebulaHandshake) XXX_Unmarshal(b []byte) error {
|
||||
return m.Unmarshal(b)
|
||||
}
|
||||
func (m *NebulaHandshake) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) {
|
||||
if deterministic {
|
||||
return xxx_messageInfo_NebulaHandshake.Marshal(b, m, deterministic)
|
||||
} else {
|
||||
b = b[:cap(b)]
|
||||
n, err := m.MarshalToSizedBuffer(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return b[:n], nil
|
||||
}
|
||||
}
|
||||
func (m *NebulaHandshake) XXX_Merge(src proto.Message) {
|
||||
xxx_messageInfo_NebulaHandshake.Merge(m, src)
|
||||
}
|
||||
func (m *NebulaHandshake) XXX_Size() int {
|
||||
return m.Size()
|
||||
}
|
||||
func (m *NebulaHandshake) XXX_DiscardUnknown() {
|
||||
xxx_messageInfo_NebulaHandshake.DiscardUnknown(m)
|
||||
}
|
||||
|
||||
var xxx_messageInfo_NebulaHandshake proto.InternalMessageInfo
|
||||
|
||||
func (m *NebulaHandshake) GetDetails() *NebulaHandshakeDetails {
|
||||
if m != nil {
|
||||
return m.Details
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *NebulaHandshake) GetHmac() []byte {
|
||||
if m != nil {
|
||||
return m.Hmac
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type NebulaHandshakeDetails struct {
|
||||
Cert []byte `protobuf:"bytes,1,opt,name=Cert,proto3" json:"Cert,omitempty"`
|
||||
InitiatorIndex uint32 `protobuf:"varint,2,opt,name=InitiatorIndex,proto3" json:"InitiatorIndex,omitempty"`
|
||||
ResponderIndex uint32 `protobuf:"varint,3,opt,name=ResponderIndex,proto3" json:"ResponderIndex,omitempty"`
|
||||
Cookie uint64 `protobuf:"varint,4,opt,name=Cookie,proto3" json:"Cookie,omitempty"`
|
||||
Time uint64 `protobuf:"varint,5,opt,name=Time,proto3" json:"Time,omitempty"`
|
||||
CertVersion uint32 `protobuf:"varint,8,opt,name=CertVersion,proto3" json:"CertVersion,omitempty"`
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) Reset() { *m = NebulaHandshakeDetails{} }
|
||||
func (m *NebulaHandshakeDetails) String() string { return proto.CompactTextString(m) }
|
||||
func (*NebulaHandshakeDetails) ProtoMessage() {}
|
||||
func (*NebulaHandshakeDetails) Descriptor() ([]byte, []int) {
|
||||
return fileDescriptor_2d65afa7693df5ef, []int{7}
|
||||
}
|
||||
func (m *NebulaHandshakeDetails) XXX_Unmarshal(b []byte) error {
|
||||
return m.Unmarshal(b)
|
||||
}
|
||||
func (m *NebulaHandshakeDetails) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) {
|
||||
if deterministic {
|
||||
return xxx_messageInfo_NebulaHandshakeDetails.Marshal(b, m, deterministic)
|
||||
} else {
|
||||
b = b[:cap(b)]
|
||||
n, err := m.MarshalToSizedBuffer(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return b[:n], nil
|
||||
}
|
||||
}
|
||||
func (m *NebulaHandshakeDetails) XXX_Merge(src proto.Message) {
|
||||
xxx_messageInfo_NebulaHandshakeDetails.Merge(m, src)
|
||||
}
|
||||
func (m *NebulaHandshakeDetails) XXX_Size() int {
|
||||
return m.Size()
|
||||
}
|
||||
func (m *NebulaHandshakeDetails) XXX_DiscardUnknown() {
|
||||
xxx_messageInfo_NebulaHandshakeDetails.DiscardUnknown(m)
|
||||
}
|
||||
|
||||
var xxx_messageInfo_NebulaHandshakeDetails proto.InternalMessageInfo
|
||||
|
||||
func (m *NebulaHandshakeDetails) GetCert() []byte {
|
||||
if m != nil {
|
||||
return m.Cert
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) GetInitiatorIndex() uint32 {
|
||||
if m != nil {
|
||||
return m.InitiatorIndex
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) GetResponderIndex() uint32 {
|
||||
if m != nil {
|
||||
return m.ResponderIndex
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) GetCookie() uint64 {
|
||||
if m != nil {
|
||||
return m.Cookie
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) GetTime() uint64 {
|
||||
if m != nil {
|
||||
return m.Time
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) GetCertVersion() uint32 {
|
||||
if m != nil {
|
||||
return m.CertVersion
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
type NebulaControl struct {
|
||||
Type NebulaControl_MessageType `protobuf:"varint,1,opt,name=Type,proto3,enum=nebula.NebulaControl_MessageType" json:"Type,omitempty"`
|
||||
InitiatorRelayIndex uint32 `protobuf:"varint,2,opt,name=InitiatorRelayIndex,proto3" json:"InitiatorRelayIndex,omitempty"`
|
||||
@@ -639,7 +503,7 @@ func (m *NebulaControl) Reset() { *m = NebulaControl{} }
|
||||
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
||||
func (*NebulaControl) ProtoMessage() {}
|
||||
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
||||
return fileDescriptor_2d65afa7693df5ef, []int{8}
|
||||
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
||||
}
|
||||
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
||||
return m.Unmarshal(b)
|
||||
@@ -729,65 +593,55 @@ func init() {
|
||||
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
||||
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
||||
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
||||
proto.RegisterType((*NebulaHandshake)(nil), "nebula.NebulaHandshake")
|
||||
proto.RegisterType((*NebulaHandshakeDetails)(nil), "nebula.NebulaHandshakeDetails")
|
||||
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
||||
}
|
||||
|
||||
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
||||
|
||||
var fileDescriptor_2d65afa7693df5ef = []byte{
|
||||
// 785 bytes of a gzipped FileDescriptorProto
|
||||
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x55, 0xcd, 0x6e, 0xeb, 0x44,
|
||||
0x14, 0x8e, 0x1d, 0x27, 0x4e, 0x4f, 0x7e, 0xae, 0x39, 0x15, 0xc1, 0x41, 0x22, 0x0a, 0x5e, 0x54,
|
||||
0x57, 0x2c, 0x72, 0x51, 0x5a, 0xae, 0x58, 0x72, 0x1b, 0x84, 0xd2, 0xaa, 0x3f, 0x61, 0x54, 0x8a,
|
||||
0xc4, 0x06, 0xb9, 0xf6, 0xd0, 0x58, 0x71, 0x3c, 0xa9, 0x3d, 0x41, 0xcd, 0x5b, 0xf0, 0x30, 0x3c,
|
||||
0x04, 0xec, 0xba, 0x42, 0x2c, 0x51, 0xbb, 0x64, 0xc9, 0x0b, 0xa0, 0x19, 0xff, 0x27, 0x86, 0xbb,
|
||||
0x9b, 0x73, 0xbe, 0xef, 0x3b, 0x73, 0xe6, 0xf3, 0x9c, 0x31, 0x74, 0x02, 0x7a, 0xb7, 0xf1, 0xed,
|
||||
0xf1, 0x3a, 0x64, 0x9c, 0x61, 0x33, 0x8e, 0xac, 0xbf, 0x55, 0x80, 0x2b, 0xb9, 0xbc, 0xa4, 0xdc,
|
||||
0xc6, 0x09, 0x68, 0x37, 0xdb, 0x35, 0x35, 0x95, 0x91, 0xf2, 0xba, 0x37, 0x19, 0x8e, 0x13, 0x4d,
|
||||
0xce, 0x18, 0x5f, 0xd2, 0x28, 0xb2, 0xef, 0xa9, 0x60, 0x11, 0xc9, 0xc5, 0x63, 0xd0, 0xbf, 0xa6,
|
||||
0xdc, 0xf6, 0xfc, 0xc8, 0x54, 0x47, 0xca, 0xeb, 0xf6, 0x64, 0xb0, 0x2f, 0x4b, 0x08, 0x24, 0x65,
|
||||
0x5a, 0xff, 0x28, 0xd0, 0x2e, 0x94, 0xc2, 0x16, 0x68, 0x57, 0x2c, 0xa0, 0x46, 0x0d, 0xbb, 0x70,
|
||||
0x30, 0x63, 0x11, 0xff, 0x76, 0x43, 0xc3, 0xad, 0xa1, 0x20, 0x42, 0x2f, 0x0b, 0x09, 0x5d, 0xfb,
|
||||
0x5b, 0x43, 0xc5, 0x8f, 0xa1, 0x2f, 0x72, 0xdf, 0xad, 0x5d, 0x9b, 0xd3, 0x2b, 0xc6, 0xbd, 0x9f,
|
||||
0x3c, 0xc7, 0xe6, 0x1e, 0x0b, 0x8c, 0x3a, 0x0e, 0xe0, 0x43, 0x81, 0x5d, 0xb2, 0x9f, 0xa9, 0x5b,
|
||||
0x82, 0xb4, 0x14, 0x9a, 0x6f, 0x02, 0x67, 0x51, 0x82, 0x1a, 0xd8, 0x03, 0x10, 0xd0, 0xf7, 0x0b,
|
||||
0x66, 0xaf, 0x3c, 0xa3, 0x89, 0x87, 0xf0, 0x2a, 0x8f, 0xe3, 0x6d, 0x75, 0xd1, 0xd9, 0xdc, 0xe6,
|
||||
0x8b, 0xe9, 0x82, 0x3a, 0x4b, 0xa3, 0x25, 0x3a, 0xcb, 0xc2, 0x98, 0x72, 0x80, 0x9f, 0xc0, 0xa0,
|
||||
0xba, 0xb3, 0x77, 0xce, 0xd2, 0x00, 0xeb, 0x77, 0x15, 0x3e, 0xd8, 0x33, 0x05, 0x2d, 0x80, 0x6b,
|
||||
0xdf, 0xbd, 0x5d, 0x07, 0xef, 0x5c, 0x37, 0x94, 0xd6, 0x77, 0x4f, 0x55, 0x53, 0x21, 0x85, 0x2c,
|
||||
0x1e, 0x81, 0x9e, 0x12, 0x9a, 0xd2, 0xe4, 0x4e, 0x6a, 0xb2, 0xc8, 0x91, 0x14, 0xc4, 0x31, 0x18,
|
||||
0xd7, 0xbe, 0x4b, 0xa8, 0x6f, 0x6f, 0x93, 0x54, 0x64, 0x36, 0x46, 0xf5, 0xa4, 0xe2, 0x1e, 0x86,
|
||||
0x13, 0xe8, 0x96, 0xc9, 0xfa, 0xa8, 0xbe, 0x57, 0xbd, 0x4c, 0xc1, 0x13, 0x68, 0xdf, 0x9e, 0x88,
|
||||
0xe5, 0x9c, 0x85, 0x5c, 0x7c, 0x74, 0xa1, 0xc0, 0x54, 0x91, 0x43, 0xa4, 0x48, 0x93, 0xaa, 0xb7,
|
||||
0xb9, 0x4a, 0xdb, 0x51, 0xbd, 0x2d, 0xa8, 0x72, 0x1a, 0x9a, 0xa0, 0x3b, 0x6c, 0x13, 0x70, 0x1a,
|
||||
0x9a, 0x75, 0x61, 0x0c, 0x49, 0x43, 0xeb, 0x08, 0x34, 0x79, 0xe2, 0x1e, 0xa8, 0x33, 0x4f, 0xba,
|
||||
0xa6, 0x11, 0x75, 0xe6, 0x89, 0xf8, 0x82, 0xc9, 0x9b, 0xa8, 0x11, 0xf5, 0x82, 0x59, 0x27, 0x00,
|
||||
0x79, 0x1b, 0x88, 0xb1, 0x2a, 0x76, 0x99, 0xc4, 0x15, 0x10, 0x34, 0x81, 0x49, 0x4d, 0x97, 0xc8,
|
||||
0xb5, 0xf5, 0x15, 0x40, 0xde, 0xc6, 0xfb, 0xf6, 0xc8, 0x2a, 0xd4, 0x0b, 0x15, 0x1e, 0xd3, 0xc1,
|
||||
0x9a, 0x7b, 0xc1, 0xfd, 0xff, 0x0f, 0x96, 0x60, 0x54, 0x0c, 0x16, 0x82, 0x76, 0xe3, 0xad, 0x68,
|
||||
0xb2, 0x8f, 0x5c, 0x5b, 0xd6, 0xde, 0xd8, 0x08, 0xb1, 0x51, 0xc3, 0x03, 0x68, 0xc4, 0x97, 0x50,
|
||||
0xb1, 0x7e, 0x84, 0x57, 0x71, 0xdd, 0x99, 0x1d, 0xb8, 0xd1, 0xc2, 0x5e, 0x52, 0xfc, 0x32, 0x9f,
|
||||
0x51, 0x45, 0x5e, 0x9f, 0x9d, 0x0e, 0x32, 0xe6, 0xee, 0xa0, 0x8a, 0x26, 0x66, 0x2b, 0xdb, 0x91,
|
||||
0x4d, 0x74, 0x88, 0x5c, 0x5b, 0x7f, 0x28, 0xd0, 0xaf, 0xd6, 0x09, 0xfa, 0x94, 0x86, 0x5c, 0xee,
|
||||
0xd2, 0x21, 0x72, 0x8d, 0x47, 0xd0, 0x3b, 0x0b, 0x3c, 0xee, 0xd9, 0x9c, 0x85, 0x67, 0x81, 0x4b,
|
||||
0x1f, 0x13, 0xa7, 0x77, 0xb2, 0x82, 0x47, 0x68, 0xb4, 0x66, 0x81, 0x4b, 0x13, 0x5e, 0xec, 0xe7,
|
||||
0x4e, 0x16, 0xfb, 0xd0, 0x9c, 0x32, 0xb6, 0xf4, 0xa8, 0xa9, 0x49, 0x67, 0x92, 0x28, 0xf3, 0xab,
|
||||
0x91, 0xfb, 0x85, 0x23, 0x68, 0x8b, 0x1e, 0x6e, 0x69, 0x18, 0x79, 0x2c, 0x30, 0x5b, 0xb2, 0x60,
|
||||
0x31, 0x75, 0xae, 0xb5, 0x9a, 0x86, 0x7e, 0xae, 0xb5, 0x74, 0xa3, 0x65, 0xfd, 0x5a, 0x87, 0x6e,
|
||||
0x7c, 0xb0, 0x29, 0x0b, 0x78, 0xc8, 0x7c, 0xfc, 0xa2, 0xf4, 0xdd, 0x3e, 0x2d, 0xbb, 0x96, 0x90,
|
||||
0x2a, 0x3e, 0xdd, 0xe7, 0x70, 0x98, 0x1d, 0x4e, 0x0e, 0x4f, 0xf1, 0xdc, 0x55, 0x90, 0x50, 0x64,
|
||||
0xc7, 0x2c, 0x28, 0x62, 0x07, 0xaa, 0x20, 0xfc, 0x0c, 0x7a, 0xe9, 0x38, 0xdf, 0x30, 0x79, 0xa9,
|
||||
0xb5, 0xec, 0xe9, 0xd8, 0x41, 0x8a, 0xcf, 0xc2, 0x37, 0x21, 0x5b, 0x49, 0x76, 0x23, 0x63, 0xef,
|
||||
0x61, 0x38, 0x86, 0x76, 0xb1, 0x70, 0xd5, 0x93, 0x53, 0x24, 0x64, 0xcf, 0x48, 0x56, 0x5c, 0xaf,
|
||||
0x50, 0x94, 0x29, 0xd6, 0xec, 0xbf, 0xfe, 0x00, 0x7d, 0xc0, 0x69, 0x48, 0x6d, 0x4e, 0x25, 0x9f,
|
||||
0xd0, 0x87, 0x0d, 0x8d, 0xb8, 0xa1, 0xe0, 0x47, 0x70, 0x58, 0xca, 0x0b, 0x4b, 0x22, 0x6a, 0xa8,
|
||||
0xa7, 0xc7, 0xbf, 0x3d, 0x0f, 0x95, 0xa7, 0xe7, 0xa1, 0xf2, 0xd7, 0xf3, 0x50, 0xf9, 0xe5, 0x65,
|
||||
0x58, 0x7b, 0x7a, 0x19, 0xd6, 0xfe, 0x7c, 0x19, 0xd6, 0x7e, 0x18, 0xdc, 0x7b, 0x7c, 0xb1, 0xb9,
|
||||
0x1b, 0x3b, 0x6c, 0xf5, 0x26, 0xf2, 0x6d, 0x67, 0xb9, 0x78, 0x78, 0x13, 0xb7, 0x74, 0xd7, 0x94,
|
||||
0x3f, 0xc2, 0xe3, 0x7f, 0x03, 0x00, 0x00, 0xff, 0xff, 0xea, 0x6f, 0xbc, 0x50, 0x18, 0x07, 0x00,
|
||||
0x00,
|
||||
// 665 bytes of a gzipped FileDescriptorProto
|
||||
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x54, 0xcd, 0x6e, 0xd3, 0x5c,
|
||||
0x10, 0x8d, 0x1d, 0x27, 0x69, 0x27, 0x4d, 0x3e, 0x7f, 0x53, 0x51, 0x12, 0x24, 0xac, 0xe0, 0x45,
|
||||
0x55, 0xb1, 0x48, 0x51, 0x5a, 0xba, 0xa6, 0x2d, 0x42, 0xa9, 0xd4, 0x9f, 0x70, 0x55, 0x8a, 0xc4,
|
||||
0xce, 0xb5, 0x2f, 0x8d, 0x55, 0xc7, 0x37, 0xb5, 0x6f, 0x50, 0xf3, 0x16, 0x3c, 0x0c, 0x0f, 0x01,
|
||||
0xbb, 0x2e, 0x59, 0xa2, 0x66, 0xc9, 0x92, 0x17, 0x40, 0xf7, 0xfa, 0xbf, 0x31, 0xb0, 0xbb, 0x33,
|
||||
0xe7, 0x9c, 0x99, 0xc9, 0xc9, 0x8c, 0x61, 0xcd, 0xa7, 0x97, 0x33, 0xcf, 0xea, 0x4f, 0x03, 0xc6,
|
||||
0x19, 0xd6, 0xa3, 0xc8, 0xfc, 0xa9, 0x02, 0x9c, 0xca, 0xe7, 0x09, 0xe5, 0x16, 0x0e, 0x40, 0x3b,
|
||||
0x9f, 0x4f, 0x69, 0x47, 0xe9, 0x29, 0x5b, 0xed, 0x81, 0xd1, 0x8f, 0x35, 0x19, 0xa3, 0x7f, 0x42,
|
||||
0xc3, 0xd0, 0xba, 0xa2, 0x82, 0x45, 0x24, 0x17, 0x77, 0xa0, 0xf1, 0x9a, 0x72, 0xcb, 0xf5, 0xc2,
|
||||
0x8e, 0xda, 0x53, 0xb6, 0x9a, 0x83, 0xee, 0xb2, 0x2c, 0x26, 0x90, 0x84, 0x69, 0xfe, 0x52, 0xa0,
|
||||
0x99, 0x2b, 0x85, 0x2b, 0xa0, 0x9d, 0x32, 0x9f, 0xea, 0x15, 0x6c, 0xc1, 0xea, 0x90, 0x85, 0xfc,
|
||||
0xed, 0x8c, 0x06, 0x73, 0x5d, 0x41, 0x84, 0x76, 0x1a, 0x12, 0x3a, 0xf5, 0xe6, 0xba, 0x8a, 0x4f,
|
||||
0x60, 0x43, 0xe4, 0xde, 0x4d, 0x1d, 0x8b, 0xd3, 0x53, 0xc6, 0xdd, 0x8f, 0xae, 0x6d, 0x71, 0x97,
|
||||
0xf9, 0x7a, 0x15, 0xbb, 0xf0, 0x48, 0x60, 0x27, 0xec, 0x13, 0x75, 0x0a, 0x90, 0x96, 0x40, 0xa3,
|
||||
0x99, 0x6f, 0x8f, 0x0b, 0x50, 0x0d, 0xdb, 0x00, 0x02, 0x7a, 0x3f, 0x66, 0xd6, 0xc4, 0xd5, 0xeb,
|
||||
0xb8, 0x0e, 0xff, 0x65, 0x71, 0xd4, 0xb6, 0x21, 0x26, 0x1b, 0x59, 0x7c, 0x7c, 0x38, 0xa6, 0xf6,
|
||||
0xb5, 0xbe, 0x22, 0x26, 0x4b, 0xc3, 0x88, 0xb2, 0x8a, 0x4f, 0xa1, 0x5b, 0x3e, 0xd9, 0xbe, 0x7d,
|
||||
0xad, 0x83, 0xf9, 0x4d, 0x85, 0xff, 0x97, 0x4c, 0x41, 0x13, 0xe0, 0xcc, 0x73, 0x2e, 0xa6, 0xfe,
|
||||
0xbe, 0xe3, 0x04, 0xd2, 0xfa, 0xd6, 0x81, 0xda, 0x51, 0x48, 0x2e, 0x8b, 0x9b, 0xd0, 0x48, 0x08,
|
||||
0x75, 0x69, 0xf2, 0x5a, 0x62, 0xb2, 0xc8, 0x91, 0x04, 0xc4, 0x3e, 0xe8, 0x67, 0x9e, 0x43, 0xa8,
|
||||
0x67, 0xcd, 0xe3, 0x54, 0xd8, 0xa9, 0xf5, 0xaa, 0x71, 0xc5, 0x25, 0x0c, 0x07, 0xd0, 0x2a, 0x92,
|
||||
0x1b, 0xbd, 0xea, 0x52, 0xf5, 0x22, 0x05, 0x77, 0xa1, 0x79, 0xb1, 0x2b, 0x9e, 0x23, 0x16, 0x70,
|
||||
0xf1, 0xa7, 0x0b, 0x05, 0x26, 0x8a, 0x0c, 0x22, 0x79, 0x9a, 0x54, 0xed, 0x65, 0x2a, 0xed, 0x81,
|
||||
0x6a, 0x2f, 0xa7, 0xca, 0x68, 0xd8, 0x81, 0x86, 0xcd, 0x66, 0x3e, 0xa7, 0x41, 0xa7, 0x2a, 0x8c,
|
||||
0x21, 0x49, 0x68, 0x6e, 0x82, 0x26, 0x7f, 0x71, 0x1b, 0xd4, 0xa1, 0x2b, 0x5d, 0xd3, 0x88, 0x3a,
|
||||
0x74, 0x45, 0x7c, 0xcc, 0xe4, 0x26, 0x6a, 0x44, 0x3d, 0x66, 0xe6, 0x2e, 0x40, 0x36, 0x06, 0x62,
|
||||
0xa4, 0x8a, 0x5c, 0x26, 0x51, 0x05, 0x04, 0x4d, 0x60, 0x52, 0xd3, 0x22, 0xf2, 0x6d, 0xbe, 0x02,
|
||||
0xc8, 0xc6, 0xf8, 0x57, 0x8f, 0xb4, 0x42, 0x35, 0x57, 0xe1, 0x36, 0x39, 0xac, 0x91, 0xeb, 0x5f,
|
||||
0xfd, 0xfd, 0xb0, 0x04, 0xa3, 0xe4, 0xb0, 0x10, 0xb4, 0x73, 0x77, 0x42, 0xe3, 0x3e, 0xf2, 0x6d,
|
||||
0x9a, 0x4b, 0x67, 0x23, 0xc4, 0x7a, 0x05, 0x57, 0xa1, 0x16, 0x2d, 0xa1, 0x62, 0x7e, 0xa9, 0x42,
|
||||
0x2b, 0x2a, 0x7c, 0xc8, 0x7c, 0x1e, 0x30, 0x0f, 0x5f, 0x16, 0xba, 0x3f, 0x2b, 0x76, 0x8f, 0x49,
|
||||
0x25, 0x03, 0xbc, 0x80, 0xf5, 0x23, 0xdf, 0xe5, 0xae, 0xc5, 0x59, 0x20, 0x57, 0xe0, 0xc8, 0x77,
|
||||
0xe8, 0x6d, 0xec, 0x53, 0x19, 0x24, 0x14, 0x84, 0x86, 0x53, 0xe6, 0x3b, 0x34, 0xaf, 0x88, 0x7c,
|
||||
0x29, 0x83, 0xf0, 0x39, 0xb4, 0x93, 0xa5, 0x3c, 0x67, 0xf2, 0xaf, 0xd1, 0xd2, 0x03, 0x78, 0x80,
|
||||
0xe4, 0x97, 0xfb, 0x4d, 0xc0, 0x26, 0x92, 0x5d, 0x4b, 0xd9, 0x4b, 0x18, 0xf6, 0xa1, 0x99, 0x2f,
|
||||
0x5c, 0x76, 0x38, 0x79, 0x42, 0x7a, 0x0c, 0x69, 0xf1, 0x46, 0x89, 0xa2, 0x48, 0x31, 0x87, 0x7f,
|
||||
0xfa, 0x8e, 0x6d, 0x00, 0x1e, 0x06, 0xd4, 0xe2, 0x54, 0xf2, 0x09, 0xbd, 0x99, 0xd1, 0x90, 0xeb,
|
||||
0x0a, 0x3e, 0x86, 0xf5, 0x42, 0x5e, 0x58, 0x12, 0x52, 0x5d, 0x3d, 0xd8, 0xf9, 0x7a, 0x6f, 0x28,
|
||||
0x77, 0xf7, 0x86, 0xf2, 0xe3, 0xde, 0x50, 0x3e, 0x2f, 0x8c, 0xca, 0xdd, 0xc2, 0xa8, 0x7c, 0x5f,
|
||||
0x18, 0x95, 0x0f, 0xdd, 0x2b, 0x97, 0x8f, 0x67, 0x97, 0x7d, 0x9b, 0x4d, 0xb6, 0x43, 0xcf, 0xb2,
|
||||
0xaf, 0xc7, 0x37, 0xdb, 0xd1, 0x48, 0x97, 0x75, 0xf9, 0x39, 0xdf, 0xf9, 0x1d, 0x00, 0x00, 0xff,
|
||||
0xff, 0x51, 0x0a, 0xe3, 0xd7, 0xde, 0x05, 0x00, 0x00,
|
||||
}
|
||||
|
||||
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
||||
@@ -1072,103 +926,6 @@ func (m *NebulaPing) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
||||
return len(dAtA) - i, nil
|
||||
}
|
||||
|
||||
func (m *NebulaHandshake) Marshal() (dAtA []byte, err error) {
|
||||
size := m.Size()
|
||||
dAtA = make([]byte, size)
|
||||
n, err := m.MarshalToSizedBuffer(dAtA[:size])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return dAtA[:n], nil
|
||||
}
|
||||
|
||||
func (m *NebulaHandshake) MarshalTo(dAtA []byte) (int, error) {
|
||||
size := m.Size()
|
||||
return m.MarshalToSizedBuffer(dAtA[:size])
|
||||
}
|
||||
|
||||
func (m *NebulaHandshake) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
||||
i := len(dAtA)
|
||||
_ = i
|
||||
var l int
|
||||
_ = l
|
||||
if len(m.Hmac) > 0 {
|
||||
i -= len(m.Hmac)
|
||||
copy(dAtA[i:], m.Hmac)
|
||||
i = encodeVarintNebula(dAtA, i, uint64(len(m.Hmac)))
|
||||
i--
|
||||
dAtA[i] = 0x12
|
||||
}
|
||||
if m.Details != nil {
|
||||
{
|
||||
size, err := m.Details.MarshalToSizedBuffer(dAtA[:i])
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
i -= size
|
||||
i = encodeVarintNebula(dAtA, i, uint64(size))
|
||||
}
|
||||
i--
|
||||
dAtA[i] = 0xa
|
||||
}
|
||||
return len(dAtA) - i, nil
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) Marshal() (dAtA []byte, err error) {
|
||||
size := m.Size()
|
||||
dAtA = make([]byte, size)
|
||||
n, err := m.MarshalToSizedBuffer(dAtA[:size])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return dAtA[:n], nil
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) MarshalTo(dAtA []byte) (int, error) {
|
||||
size := m.Size()
|
||||
return m.MarshalToSizedBuffer(dAtA[:size])
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
||||
i := len(dAtA)
|
||||
_ = i
|
||||
var l int
|
||||
_ = l
|
||||
if m.CertVersion != 0 {
|
||||
i = encodeVarintNebula(dAtA, i, uint64(m.CertVersion))
|
||||
i--
|
||||
dAtA[i] = 0x40
|
||||
}
|
||||
if m.Time != 0 {
|
||||
i = encodeVarintNebula(dAtA, i, uint64(m.Time))
|
||||
i--
|
||||
dAtA[i] = 0x28
|
||||
}
|
||||
if m.Cookie != 0 {
|
||||
i = encodeVarintNebula(dAtA, i, uint64(m.Cookie))
|
||||
i--
|
||||
dAtA[i] = 0x20
|
||||
}
|
||||
if m.ResponderIndex != 0 {
|
||||
i = encodeVarintNebula(dAtA, i, uint64(m.ResponderIndex))
|
||||
i--
|
||||
dAtA[i] = 0x18
|
||||
}
|
||||
if m.InitiatorIndex != 0 {
|
||||
i = encodeVarintNebula(dAtA, i, uint64(m.InitiatorIndex))
|
||||
i--
|
||||
dAtA[i] = 0x10
|
||||
}
|
||||
if len(m.Cert) > 0 {
|
||||
i -= len(m.Cert)
|
||||
copy(dAtA[i:], m.Cert)
|
||||
i = encodeVarintNebula(dAtA, i, uint64(len(m.Cert)))
|
||||
i--
|
||||
dAtA[i] = 0xa
|
||||
}
|
||||
return len(dAtA) - i, nil
|
||||
}
|
||||
|
||||
func (m *NebulaControl) Marshal() (dAtA []byte, err error) {
|
||||
size := m.Size()
|
||||
dAtA = make([]byte, size)
|
||||
@@ -1375,51 +1132,6 @@ func (m *NebulaPing) Size() (n int) {
|
||||
return n
|
||||
}
|
||||
|
||||
func (m *NebulaHandshake) Size() (n int) {
|
||||
if m == nil {
|
||||
return 0
|
||||
}
|
||||
var l int
|
||||
_ = l
|
||||
if m.Details != nil {
|
||||
l = m.Details.Size()
|
||||
n += 1 + l + sovNebula(uint64(l))
|
||||
}
|
||||
l = len(m.Hmac)
|
||||
if l > 0 {
|
||||
n += 1 + l + sovNebula(uint64(l))
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func (m *NebulaHandshakeDetails) Size() (n int) {
|
||||
if m == nil {
|
||||
return 0
|
||||
}
|
||||
var l int
|
||||
_ = l
|
||||
l = len(m.Cert)
|
||||
if l > 0 {
|
||||
n += 1 + l + sovNebula(uint64(l))
|
||||
}
|
||||
if m.InitiatorIndex != 0 {
|
||||
n += 1 + sovNebula(uint64(m.InitiatorIndex))
|
||||
}
|
||||
if m.ResponderIndex != 0 {
|
||||
n += 1 + sovNebula(uint64(m.ResponderIndex))
|
||||
}
|
||||
if m.Cookie != 0 {
|
||||
n += 1 + sovNebula(uint64(m.Cookie))
|
||||
}
|
||||
if m.Time != 0 {
|
||||
n += 1 + sovNebula(uint64(m.Time))
|
||||
}
|
||||
if m.CertVersion != 0 {
|
||||
n += 1 + sovNebula(uint64(m.CertVersion))
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func (m *NebulaControl) Size() (n int) {
|
||||
if m == nil {
|
||||
return 0
|
||||
@@ -2236,305 +1948,6 @@ func (m *NebulaPing) Unmarshal(dAtA []byte) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (m *NebulaHandshake) Unmarshal(dAtA []byte) error {
|
||||
l := len(dAtA)
|
||||
iNdEx := 0
|
||||
for iNdEx < l {
|
||||
preIndex := iNdEx
|
||||
var wire uint64
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
wire |= uint64(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
fieldNum := int32(wire >> 3)
|
||||
wireType := int(wire & 0x7)
|
||||
if wireType == 4 {
|
||||
return fmt.Errorf("proto: NebulaHandshake: wiretype end group for non-group")
|
||||
}
|
||||
if fieldNum <= 0 {
|
||||
return fmt.Errorf("proto: NebulaHandshake: illegal tag %d (wire type %d)", fieldNum, wire)
|
||||
}
|
||||
switch fieldNum {
|
||||
case 1:
|
||||
if wireType != 2 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field Details", wireType)
|
||||
}
|
||||
var msglen int
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
msglen |= int(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
if msglen < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
postIndex := iNdEx + msglen
|
||||
if postIndex < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
if postIndex > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
if m.Details == nil {
|
||||
m.Details = &NebulaHandshakeDetails{}
|
||||
}
|
||||
if err := m.Details.Unmarshal(dAtA[iNdEx:postIndex]); err != nil {
|
||||
return err
|
||||
}
|
||||
iNdEx = postIndex
|
||||
case 2:
|
||||
if wireType != 2 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field Hmac", wireType)
|
||||
}
|
||||
var byteLen int
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
byteLen |= int(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
if byteLen < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
postIndex := iNdEx + byteLen
|
||||
if postIndex < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
if postIndex > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
m.Hmac = append(m.Hmac[:0], dAtA[iNdEx:postIndex]...)
|
||||
if m.Hmac == nil {
|
||||
m.Hmac = []byte{}
|
||||
}
|
||||
iNdEx = postIndex
|
||||
default:
|
||||
iNdEx = preIndex
|
||||
skippy, err := skipNebula(dAtA[iNdEx:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if (skippy < 0) || (iNdEx+skippy) < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
if (iNdEx + skippy) > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
iNdEx += skippy
|
||||
}
|
||||
}
|
||||
|
||||
if iNdEx > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (m *NebulaHandshakeDetails) Unmarshal(dAtA []byte) error {
|
||||
l := len(dAtA)
|
||||
iNdEx := 0
|
||||
for iNdEx < l {
|
||||
preIndex := iNdEx
|
||||
var wire uint64
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
wire |= uint64(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
fieldNum := int32(wire >> 3)
|
||||
wireType := int(wire & 0x7)
|
||||
if wireType == 4 {
|
||||
return fmt.Errorf("proto: NebulaHandshakeDetails: wiretype end group for non-group")
|
||||
}
|
||||
if fieldNum <= 0 {
|
||||
return fmt.Errorf("proto: NebulaHandshakeDetails: illegal tag %d (wire type %d)", fieldNum, wire)
|
||||
}
|
||||
switch fieldNum {
|
||||
case 1:
|
||||
if wireType != 2 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field Cert", wireType)
|
||||
}
|
||||
var byteLen int
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
byteLen |= int(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
if byteLen < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
postIndex := iNdEx + byteLen
|
||||
if postIndex < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
if postIndex > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
m.Cert = append(m.Cert[:0], dAtA[iNdEx:postIndex]...)
|
||||
if m.Cert == nil {
|
||||
m.Cert = []byte{}
|
||||
}
|
||||
iNdEx = postIndex
|
||||
case 2:
|
||||
if wireType != 0 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field InitiatorIndex", wireType)
|
||||
}
|
||||
m.InitiatorIndex = 0
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
m.InitiatorIndex |= uint32(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
case 3:
|
||||
if wireType != 0 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field ResponderIndex", wireType)
|
||||
}
|
||||
m.ResponderIndex = 0
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
m.ResponderIndex |= uint32(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
case 4:
|
||||
if wireType != 0 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field Cookie", wireType)
|
||||
}
|
||||
m.Cookie = 0
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
m.Cookie |= uint64(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
case 5:
|
||||
if wireType != 0 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field Time", wireType)
|
||||
}
|
||||
m.Time = 0
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
m.Time |= uint64(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
case 8:
|
||||
if wireType != 0 {
|
||||
return fmt.Errorf("proto: wrong wireType = %d for field CertVersion", wireType)
|
||||
}
|
||||
m.CertVersion = 0
|
||||
for shift := uint(0); ; shift += 7 {
|
||||
if shift >= 64 {
|
||||
return ErrIntOverflowNebula
|
||||
}
|
||||
if iNdEx >= l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
b := dAtA[iNdEx]
|
||||
iNdEx++
|
||||
m.CertVersion |= uint32(b&0x7F) << shift
|
||||
if b < 0x80 {
|
||||
break
|
||||
}
|
||||
}
|
||||
default:
|
||||
iNdEx = preIndex
|
||||
skippy, err := skipNebula(dAtA[iNdEx:])
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if (skippy < 0) || (iNdEx+skippy) < 0 {
|
||||
return ErrInvalidLengthNebula
|
||||
}
|
||||
if (iNdEx + skippy) > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
iNdEx += skippy
|
||||
}
|
||||
}
|
||||
|
||||
if iNdEx > l {
|
||||
return io.ErrUnexpectedEOF
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (m *NebulaControl) Unmarshal(dAtA []byte) error {
|
||||
l := len(dAtA)
|
||||
iNdEx := 0
|
||||
|
||||
+3
-15
@@ -60,21 +60,9 @@ message NebulaPing {
|
||||
uint64 Time = 2;
|
||||
}
|
||||
|
||||
message NebulaHandshake {
|
||||
NebulaHandshakeDetails Details = 1;
|
||||
bytes Hmac = 2;
|
||||
}
|
||||
|
||||
message NebulaHandshakeDetails {
|
||||
bytes Cert = 1;
|
||||
uint32 InitiatorIndex = 2;
|
||||
uint32 ResponderIndex = 3;
|
||||
uint64 Cookie = 4;
|
||||
uint64 Time = 5;
|
||||
uint32 CertVersion = 8;
|
||||
// reserved for WIP multiport
|
||||
reserved 6, 7;
|
||||
}
|
||||
// NebulaHandshake / NebulaHandshakeDetails moved to
|
||||
// handshake/handshake.proto. The handshake package speaks that wire format
|
||||
// directly via a hand-written encoder/decoder.
|
||||
|
||||
message NebulaControl {
|
||||
enum MessageType {
|
||||
|
||||
@@ -1,75 +0,0 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
type endianness interface {
|
||||
PutUint64(b []byte, v uint64)
|
||||
}
|
||||
|
||||
var noiseEndianness endianness = binary.BigEndian
|
||||
|
||||
type NebulaCipherState struct {
|
||||
c noise.Cipher
|
||||
//k [32]byte
|
||||
//n uint64
|
||||
}
|
||||
|
||||
func NewNebulaCipherState(s *noise.CipherState) *NebulaCipherState {
|
||||
return &NebulaCipherState{c: s.Cipher()}
|
||||
|
||||
}
|
||||
|
||||
// EncryptDanger encrypts and authenticates a given payload.
|
||||
//
|
||||
// out is a destination slice to hold the output of the EncryptDanger operation.
|
||||
// - ad is additional data, which will be authenticated and appended to out, but not encrypted.
|
||||
// - plaintext is encrypted, authenticated and appended to out.
|
||||
// - n is a nonce value which must never be re-used with this key.
|
||||
// - nb is a buffer used for temporary storage in the implementation of this call, which should
|
||||
// be re-used by callers to minimize garbage collection.
|
||||
func (s *NebulaCipherState) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
if s != nil {
|
||||
// TODO: Is this okay now that we have made messageCounter atomic?
|
||||
// Alternative may be to split the counter space into ranges
|
||||
//if n <= s.n {
|
||||
// return nil, errors.New("CRITICAL: a duplicate counter value was used")
|
||||
//}
|
||||
//s.n = n
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
nb[3] = 0
|
||||
noiseEndianness.PutUint64(nb[4:], n)
|
||||
out = s.c.(cipher.AEAD).Seal(out, nb, plaintext, ad)
|
||||
//l.Debugf("Encryption: outlen: %d, nonce: %d, ad: %s, plainlen %d", len(out), n, ad, len(plaintext))
|
||||
return out, nil
|
||||
} else {
|
||||
return nil, errors.New("no cipher state available to encrypt")
|
||||
}
|
||||
}
|
||||
|
||||
func (s *NebulaCipherState) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
if s != nil {
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
nb[3] = 0
|
||||
noiseEndianness.PutUint64(nb[4:], n)
|
||||
return s.c.(cipher.AEAD).Open(out, nb, ciphertext, ad)
|
||||
} else {
|
||||
return []byte{}, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (s *NebulaCipherState) Overhead() int {
|
||||
if s != nil {
|
||||
return s.c.(cipher.AEAD).Overhead()
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// CipherStateAESGCM is the data-plane wrapper for the AES-GCM AEAD cipher.
|
||||
// AES-GCM uses big-endian nonce encoding per the Noise spec.
|
||||
type CipherStateAESGCM struct {
|
||||
c cipher.AEAD
|
||||
}
|
||||
|
||||
// NewCipherStateAESGCM extracts the underlying AEAD from the post-handshake noise.CipherState.
|
||||
// The caller is responsible for ensuring the noise cipher is actually AES-GCM,
|
||||
// otherwise the type assertion still succeeds but the nonce endianness will be wrong on the wire.
|
||||
func NewCipherStateAESGCM(s *noise.CipherState) *CipherStateAESGCM {
|
||||
return &CipherStateAESGCM{c: s.Cipher().(cipher.AEAD)}
|
||||
}
|
||||
|
||||
func (s *CipherStateAESGCM) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
if s == nil {
|
||||
return nil, errors.New("no cipher state available to encrypt")
|
||||
}
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
nb[3] = 0
|
||||
binary.BigEndian.PutUint64(nb[4:], n)
|
||||
return s.c.Seal(out, nb, plaintext, ad), nil
|
||||
}
|
||||
|
||||
func (s *CipherStateAESGCM) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
if s == nil {
|
||||
return []byte{}, nil
|
||||
}
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
nb[3] = 0
|
||||
binary.BigEndian.PutUint64(nb[4:], n)
|
||||
return s.c.Open(out, nb, ciphertext, ad)
|
||||
}
|
||||
|
||||
func (s *CipherStateAESGCM) Overhead() int {
|
||||
if s == nil {
|
||||
return 0
|
||||
}
|
||||
return s.c.Overhead()
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// CipherStateChaChaPoly is the data-plane wrapper for the ChaCha20-Poly1305 AEAD cipher.
|
||||
// ChaCha20-Poly1305 uses little-endian nonce encoding per the Noise spec.
|
||||
type CipherStateChaChaPoly struct {
|
||||
c cipher.AEAD
|
||||
}
|
||||
|
||||
// NewCipherStateChaChaPoly extracts the underlying AEAD from the post-handshake noise.CipherState.
|
||||
// The caller is responsible for ensuring the noise cipher is actually ChaCha20-Poly1305.
|
||||
func NewCipherStateChaChaPoly(s *noise.CipherState) *CipherStateChaChaPoly {
|
||||
return &CipherStateChaChaPoly{c: s.Cipher().(cipher.AEAD)}
|
||||
}
|
||||
|
||||
func (s *CipherStateChaChaPoly) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
if s == nil {
|
||||
return nil, errors.New("no cipher state available to encrypt")
|
||||
}
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
nb[3] = 0
|
||||
binary.LittleEndian.PutUint64(nb[4:], n)
|
||||
return s.c.Seal(out, nb, plaintext, ad), nil
|
||||
}
|
||||
|
||||
func (s *CipherStateChaChaPoly) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||
if s == nil {
|
||||
return []byte{}, nil
|
||||
}
|
||||
nb[0] = 0
|
||||
nb[1] = 0
|
||||
nb[2] = 0
|
||||
nb[3] = 0
|
||||
binary.LittleEndian.PutUint64(nb[4:], n)
|
||||
return s.c.Open(out, nb, ciphertext, ad)
|
||||
}
|
||||
|
||||
func (s *CipherStateChaChaPoly) Overhead() int {
|
||||
if s == nil {
|
||||
return 0
|
||||
}
|
||||
return s.c.Overhead()
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
)
|
||||
|
||||
// CipherState is the post-handshake AEAD cipher used for the data plane.
|
||||
// Each supported cipher has its own concrete implementation in this package with the nonce endianness hardcoded,
|
||||
// so the encrypt/decrypt fast path avoids interface dispatch on the byte order.
|
||||
type CipherState interface {
|
||||
// EncryptDanger encrypts and authenticates a given payload.
|
||||
//
|
||||
// out is a destination slice to hold the output of the EncryptDanger operation.
|
||||
// - ad is additional data, which will be authenticated and appended to out, but not encrypted.
|
||||
// - plaintext is encrypted, authenticated and appended to out.
|
||||
// - n is a nonce value which must never be re-used with this key.
|
||||
// - nb is a scratch buffer used to assemble the nonce.
|
||||
EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error)
|
||||
|
||||
// DecryptDanger authenticates and decrypts a given payload, with the same argument shape as EncryptDanger.
|
||||
DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error)
|
||||
|
||||
// Overhead returns the AEAD tag size, or 0 if the receiver is nil.
|
||||
Overhead() int
|
||||
}
|
||||
|
||||
// NewCipherState wraps the post-handshake noise.CipherState in the per-cipher type that matches cipherFunc.
|
||||
// cipherFunc must be the same cipher used to build the noise CipherSuite that produced s.
|
||||
func NewCipherState(s *noise.CipherState, cipherFunc noise.CipherFunc) CipherState {
|
||||
switch cipherFunc.CipherName() {
|
||||
case CipherAESGCM.CipherName():
|
||||
return NewCipherStateAESGCM(s)
|
||||
case noise.CipherChaChaPoly.CipherName():
|
||||
return NewCipherStateChaChaPoly(s)
|
||||
default:
|
||||
panic(fmt.Sprintf("noiseutil: unsupported cipher %q", cipherFunc.CipherName()))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,166 @@
|
||||
package noiseutil
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/flynn/noise"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestCipherStateAESGCMRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, CipherAESGCM)
|
||||
roundtrip(t, NewCipherStateAESGCM(enc), NewCipherStateAESGCM(dec))
|
||||
}
|
||||
|
||||
func TestCipherStateChaChaPolyRoundtrip(t *testing.T) {
|
||||
enc, dec := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
roundtrip(t, NewCipherStateChaChaPoly(enc), NewCipherStateChaChaPoly(dec))
|
||||
}
|
||||
|
||||
func TestNewCipherStateDispatch(t *testing.T) {
|
||||
encA, _ := buildCipherStates(t, CipherAESGCM)
|
||||
encC, _ := buildCipherStates(t, noise.CipherChaChaPoly)
|
||||
|
||||
assert.IsType(t, &CipherStateAESGCM{}, NewCipherState(encA, CipherAESGCM))
|
||||
assert.IsType(t, &CipherStateChaChaPoly{}, NewCipherState(encC, noise.CipherChaChaPoly))
|
||||
}
|
||||
|
||||
func TestNewCipherStateUnsupportedPanics(t *testing.T) {
|
||||
enc, _ := buildCipherStates(t, CipherAESGCM)
|
||||
assert.Panics(t, func() {
|
||||
NewCipherState(enc, fakeCipher{})
|
||||
})
|
||||
}
|
||||
|
||||
type fakeCipher struct{}
|
||||
|
||||
func (fakeCipher) Cipher(k [32]byte) noise.Cipher { return nil }
|
||||
func (fakeCipher) CipherName() string { return "Fake" }
|
||||
|
||||
// buildCipherStates runs an in-memory NN handshake with the requested cipher
|
||||
// to produce a pair of post-handshake CipherStates that share keys.
|
||||
func buildCipherStates(t *testing.T, c noise.CipherFunc) (*noise.CipherState, *noise.CipherState) {
|
||||
t.Helper()
|
||||
suite := noise.NewCipherSuite(noise.DH25519, c, noise.HashSHA256)
|
||||
cfg := noise.Config{CipherSuite: suite, Pattern: noise.HandshakeNN}
|
||||
cfg.Initiator = true
|
||||
hsI, err := noise.NewHandshakeState(cfg)
|
||||
require.NoError(t, err)
|
||||
cfg.Initiator = false
|
||||
hsR, err := noise.NewHandshakeState(cfg)
|
||||
require.NoError(t, err)
|
||||
|
||||
msg, _, _, err := hsI.WriteMessage(nil, nil)
|
||||
require.NoError(t, err)
|
||||
_, _, _, err = hsR.ReadMessage(nil, msg)
|
||||
require.NoError(t, err)
|
||||
|
||||
msg, dR, _, err := hsR.WriteMessage(nil, nil)
|
||||
require.NoError(t, err)
|
||||
_, eI, _, err := hsI.ReadMessage(nil, msg)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, eI)
|
||||
require.NotNil(t, dR)
|
||||
|
||||
// noise returns (cs1, cs2) where cs1 is the initiator->responder cipher.
|
||||
return eI, dR
|
||||
}
|
||||
|
||||
func roundtrip(t *testing.T, enc, dec CipherState) {
|
||||
t.Helper()
|
||||
plaintext := []byte("nebula cipher state roundtrip")
|
||||
ad := []byte("aad")
|
||||
nb := make([]byte, 12)
|
||||
|
||||
ct, err := enc.EncryptDanger(nil, ad, plaintext, 1, nb)
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, plaintext, ct)
|
||||
|
||||
pt, err := dec.DecryptDanger(nil, ad, ct, 1, nb)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, plaintext, pt)
|
||||
|
||||
// Wrong nonce must fail authentication.
|
||||
_, err = dec.DecryptDanger(nil, ad, ct, 2, nb)
|
||||
require.Error(t, err)
|
||||
|
||||
assert.Equal(t, enc.Overhead(), dec.Overhead())
|
||||
assert.Equal(t, 16, enc.Overhead())
|
||||
}
|
||||
|
||||
func BenchmarkCipherStateEncryptAESGCM(b *testing.B) {
|
||||
enc, _ := buildCipherStatesB(b, CipherAESGCM)
|
||||
benchEncryptCipherState(b, NewCipherState(enc, CipherAESGCM))
|
||||
}
|
||||
|
||||
func BenchmarkCipherStateEncryptChaChaPoly(b *testing.B) {
|
||||
enc, _ := buildCipherStatesB(b, noise.CipherChaChaPoly)
|
||||
benchEncryptCipherState(b, NewCipherState(enc, noise.CipherChaChaPoly))
|
||||
}
|
||||
|
||||
func benchEncryptCipherState(b *testing.B, cs CipherState) {
|
||||
plaintext := make([]byte, 1280)
|
||||
ad := make([]byte, 16)
|
||||
nb := make([]byte, 12)
|
||||
out := make([]byte, 0, len(plaintext)+cs.Overhead())
|
||||
b.ResetTimer()
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
var err error
|
||||
out, err = cs.EncryptDanger(out[:0], ad, plaintext, uint64(i+1), nb)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func buildCipherStatesB(b *testing.B, c noise.CipherFunc) (*noise.CipherState, *noise.CipherState) {
|
||||
b.Helper()
|
||||
suite := noise.NewCipherSuite(noise.DH25519, c, noise.HashSHA256)
|
||||
cfg := noise.Config{CipherSuite: suite, Pattern: noise.HandshakeNN}
|
||||
cfg.Initiator = true
|
||||
hsI, err := noise.NewHandshakeState(cfg)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
cfg.Initiator = false
|
||||
hsR, err := noise.NewHandshakeState(cfg)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
msg, _, _, err := hsI.WriteMessage(nil, nil)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if _, _, _, err := hsR.ReadMessage(nil, msg); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
msg, dR, _, err := hsR.WriteMessage(nil, nil)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
_, eI, _, err := hsI.ReadMessage(nil, msg)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
return eI, dR
|
||||
}
|
||||
|
||||
func TestCipherStateNilSafety(t *testing.T) {
|
||||
var aes *CipherStateAESGCM
|
||||
_, err := aes.EncryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||
require.Error(t, err)
|
||||
out, err := aes.DecryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, out)
|
||||
assert.Equal(t, 0, aes.Overhead())
|
||||
|
||||
var cc *CipherStateChaChaPoly
|
||||
_, err = cc.EncryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||
require.Error(t, err)
|
||||
out, err = cc.DecryptDanger(nil, nil, nil, 0, make([]byte, 12))
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, out)
|
||||
assert.Equal(t, 0, cc.Overhead())
|
||||
}
|
||||
+311
-426
@@ -1,233 +1,252 @@
|
||||
package nebula
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"time"
|
||||
|
||||
"github.com/google/gopacket/layers"
|
||||
"golang.org/x/net/ipv6"
|
||||
|
||||
"github.com/sirupsen/logrus"
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"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")
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||
// TODO: record metrics for rx holepunch/punchy packets?
|
||||
if len(packet) > 1 {
|
||||
f.l.WithField("packet", packet).Infof("Error while parsing inbound packet from %s: %s", via, err)
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Error while parsing inbound packet",
|
||||
"from", via,
|
||||
"error", err,
|
||||
"packet", packet,
|
||||
)
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if h.Version != header.Version {
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Unexpected header version received", "from", via)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Check before processing to see if this is a expected type/subtype
|
||||
if !h.IsValidSubType() {
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Unexpected packet received", "from", via)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
//l.Error("in packet ", header, packet[HeaderLen:])
|
||||
if !via.IsRelayed {
|
||||
if f.myVpnNetworksTable.Contains(via.UdpAddr.Addr()) {
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("from", via).Debug("Refusing to process double encrypted packet")
|
||||
f.messageMetrics.RxInvalid(1)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Refusing to process double encrypted packet", "from", via)
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// don't keep Rx metrics for message type, since you can see those in the tun metrics
|
||||
if h.Type != header.Message {
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
}
|
||||
|
||||
// Unencrypted packets
|
||||
switch h.Type {
|
||||
case header.Handshake:
|
||||
f.handshakeManager.HandleIncoming(via, packet, h)
|
||||
return
|
||||
|
||||
case header.RecvError:
|
||||
f.handleRecvError(via.UdpAddr, h)
|
||||
return
|
||||
}
|
||||
|
||||
// Relay packets are special
|
||||
isMessageRelay := (h.Type == header.Message && h.Subtype == header.MessageRelay)
|
||||
|
||||
var hostinfo *HostInfo
|
||||
// verify if we've seen this index before, otherwise respond to the handshake initiation
|
||||
if h.Type == header.Message && h.Subtype == header.MessageRelay {
|
||||
if isMessageRelay {
|
||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||
} else {
|
||||
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
||||
}
|
||||
|
||||
var ci *ConnectionState
|
||||
if hostinfo != nil {
|
||||
ci = hostinfo.ConnectionState
|
||||
// At this point we should have a valid existing tunnel, verify and send
|
||||
// recvError if necessary
|
||||
if hostinfo == nil || hostinfo.ConnectionState == nil {
|
||||
if !via.IsRelayed {
|
||||
f.maybeSendRecvError(via.UdpAddr, h.RemoteIndex)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// All remaining packets are encrypted
|
||||
ci := hostinfo.ConnectionState
|
||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
||||
return
|
||||
}
|
||||
|
||||
// Relay packets are special
|
||||
if isMessageRelay {
|
||||
f.handleOutsideRelayPacket(hostinfo, via, out, packet, h, fwPacket, parsedRx, lhf, nb, q, localCache, meta)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
out, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
if err != nil {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Failed to decrypt packet",
|
||||
"error", err,
|
||||
"from", via,
|
||||
"header", h,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Roam before we respond
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
f.connectionManager.In(hostinfo)
|
||||
|
||||
switch h.Type {
|
||||
case header.Message:
|
||||
if !f.handleEncrypted(ci, via, h) {
|
||||
return
|
||||
}
|
||||
|
||||
switch h.Subtype {
|
||||
case header.MessageNone:
|
||||
if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) {
|
||||
return
|
||||
}
|
||||
case header.MessageRelay:
|
||||
// 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 is guaranteed to be at least 16 bytes at this point, b/c it got past the h.Parse() call above. If it's
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Successfully validated the thing. Get rid of the Relay header.
|
||||
signedPayload = signedPayload[header.Len:]
|
||||
// Pull the Roaming parts up here, and return in all call paths.
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||
f.connectionManager.In(hostinfo)
|
||||
f.connectionManager.RelayUsed(h.RemoteIndex)
|
||||
|
||||
relay, ok := hostinfo.relayState.QueryRelayForByIdx(h.RemoteIndex)
|
||||
if !ok {
|
||||
// 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.
|
||||
hostinfo.logger(f.l).WithFields(logrus.Fields{"vpnAddrs": hostinfo.vpnAddrs, "remoteIndex": h.RemoteIndex}).Error("HostInfo missing remote relay index")
|
||||
return
|
||||
}
|
||||
|
||||
switch relay.Type {
|
||||
case TerminalType:
|
||||
// If I am the target of this relay, process the unwrapped packet
|
||||
// From this recursive point, all these variables are 'burned'. We shouldn't rely on them again.
|
||||
via = ViaSender{
|
||||
UdpAddr: via.UdpAddr,
|
||||
relayHI: hostinfo,
|
||||
remoteIdx: relay.RemoteIndex,
|
||||
relay: relay,
|
||||
IsRelayed: true,
|
||||
}
|
||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
||||
return
|
||||
case ForwardingType:
|
||||
// Find the target HostInfo relay object
|
||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithField("relayTo", relay.PeerAddr).WithError(err).WithField("hostinfo.vpnAddrs", hostinfo.vpnAddrs).Info("Failed to find target host info by ip")
|
||||
return
|
||||
}
|
||||
|
||||
// If that relay is Established, forward the payload through it
|
||||
if targetRelay.State == Established {
|
||||
switch targetRelay.Type {
|
||||
case ForwardingType:
|
||||
// Forward this packet through the relay tunnel
|
||||
// Find the target HostInfo
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||
return
|
||||
case TerminalType:
|
||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||
}
|
||||
} else {
|
||||
hostinfo.logger(f.l).WithFields(logrus.Fields{"relayTo": relay.PeerAddr, "relayFrom": hostinfo.vpnAddrs[0], "targetRelayState": targetRelay.State}).Info("Unexpected target relay state")
|
||||
return
|
||||
}
|
||||
}
|
||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, parsedRx, nb, q, localCache, meta)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||
return
|
||||
}
|
||||
|
||||
case header.LightHouse:
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
if !f.handleEncrypted(ci, via, h) {
|
||||
return
|
||||
}
|
||||
|
||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).WithField("from", via).
|
||||
WithField("packet", packet).
|
||||
Error("Failed to decrypt lighthouse packet")
|
||||
return
|
||||
}
|
||||
|
||||
//TODO: assert via is not relayed
|
||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, d, f)
|
||||
|
||||
// Fallthrough to the bottom to record incoming traffic
|
||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||
|
||||
case header.Test:
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
if !f.handleEncrypted(ci, via, h) {
|
||||
switch h.Subtype {
|
||||
case header.TestReply:
|
||||
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||
case header.TestRequest:
|
||||
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, out)
|
||||
default:
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||
return
|
||||
}
|
||||
|
||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).WithField("from", via).
|
||||
WithField("packet", packet).
|
||||
Error("Failed to decrypt test packet")
|
||||
return
|
||||
}
|
||||
|
||||
if h.Subtype == header.TestRequest {
|
||||
// This testRequest might be from TryPromoteBest, so we should roam
|
||||
// to the new IP address before responding
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
f.send(header.Test, header.TestReply, ci, hostinfo, d, nb, out)
|
||||
}
|
||||
|
||||
// Fallthrough to the bottom to record incoming traffic
|
||||
|
||||
// Non encrypted messages below here, they should not fall through to avoid tracking incoming traffic since they
|
||||
// are unauthenticated
|
||||
|
||||
case header.Handshake:
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
f.handshakeManager.HandleIncoming(via, packet, h)
|
||||
return
|
||||
|
||||
case header.RecvError:
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
f.handleRecvError(via.UdpAddr, h)
|
||||
return
|
||||
|
||||
case header.CloseTunnel:
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
if !f.handleEncrypted(ci, via, h) {
|
||||
return
|
||||
}
|
||||
_, err = f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).WithField("from", via).
|
||||
WithField("packet", packet).
|
||||
Error("Failed to decrypt CloseTunnel packet")
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo.logger(f.l).WithField("from", via).
|
||||
Info("Close tunnel received, tearing down.")
|
||||
|
||||
hostinfo.logger(f.l).Info("Close tunnel received, tearing down.", "from", via)
|
||||
f.closeTunnel(hostinfo)
|
||||
return
|
||||
|
||||
case header.Control:
|
||||
if !f.handleEncrypted(ci, via, h) {
|
||||
return
|
||||
}
|
||||
|
||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).WithField("from", via).
|
||||
WithField("packet", packet).
|
||||
Error("Failed to decrypt Control packet")
|
||||
return
|
||||
}
|
||||
|
||||
f.relayManager.HandleControlMsg(hostinfo, d, f)
|
||||
f.relayManager.HandleControlMsg(hostinfo, out, f)
|
||||
|
||||
default:
|
||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||
hostinfo.logger(f.l).Debugf("Unexpected packet received from %s", via)
|
||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message type seen", "from", via, "header", h)
|
||||
}
|
||||
}
|
||||
|
||||
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 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
|
||||
// otherwise malformed (meaning, there is no trailing 16 byte AEAD value), then this will result in at worst a 0-length slice
|
||||
// which will gracefully fail in the DecryptDanger call.
|
||||
signedPayload := packet[:len(packet)-hostinfo.ConnectionState.dKey.Overhead()]
|
||||
signatureValue := packet[len(packet)-hostinfo.ConnectionState.dKey.Overhead():]
|
||||
var err error
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, signedPayload, signatureValue, h.MessageCounter, nb)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Successfully validated the thing. Get rid of the Relay header.
|
||||
signedPayload = signedPayload[header.Len:]
|
||||
// Pull the Roaming parts up here, and return in all call paths.
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
// Track usage of both the HostInfo and the Relay for the received & authenticated packet
|
||||
f.connectionManager.In(hostinfo)
|
||||
f.connectionManager.RelayUsed(h.RemoteIndex)
|
||||
|
||||
relay, ok := hostinfo.relayState.QueryRelayForByIdx(h.RemoteIndex)
|
||||
if !ok {
|
||||
// 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.
|
||||
hostinfo.logger(f.l).Error("HostInfo missing remote relay index",
|
||||
"vpnAddrs", hostinfo.vpnAddrs,
|
||||
"remoteIndex", h.RemoteIndex,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
f.handleHostRoaming(hostinfo, via)
|
||||
switch relay.Type {
|
||||
case TerminalType:
|
||||
// If I am the target of this relay, process the unwrapped packet
|
||||
// From this recursive point, all these variables are 'burned'. We shouldn't rely on them again.
|
||||
via = ViaSender{
|
||||
UdpAddr: via.UdpAddr,
|
||||
relayHI: hostinfo,
|
||||
remoteIdx: relay.RemoteIndex,
|
||||
relay: relay,
|
||||
IsRelayed: true,
|
||||
}
|
||||
f.readOutsidePackets(via, out[:0], signedPayload, h, fwPacket, parsedRx, lhf, nb, q, localCache, meta)
|
||||
return
|
||||
case ForwardingType:
|
||||
// Find the target HostInfo relay object
|
||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
||||
"relayTo", relay.PeerAddr,
|
||||
"error", err,
|
||||
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
f.connectionManager.In(hostinfo)
|
||||
// If that relay is Established, forward the payload through it
|
||||
if targetRelay.State == Established {
|
||||
switch targetRelay.Type {
|
||||
case ForwardingType:
|
||||
// Forward this packet through the relay tunnel
|
||||
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||
case TerminalType:
|
||||
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||
return
|
||||
default:
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Unexpected targetRelay Type", "from", via, "relayType", targetRelay.Type)
|
||||
}
|
||||
return
|
||||
}
|
||||
} else {
|
||||
hostinfo.logger(f.l).Info("Unexpected target relay state",
|
||||
"relayTo", relay.PeerAddr,
|
||||
"relayFrom", hostinfo.vpnAddrs[0],
|
||||
"targetRelayState", targetRelay.State,
|
||||
)
|
||||
return
|
||||
}
|
||||
default:
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Unexpected relay type", "from", via, "relayType", relay.Type)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
||||
@@ -247,20 +266,27 @@ func (f *Interface) sendCloseTunnel(h *HostInfo) {
|
||||
func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
||||
if !via.IsRelayed && hostinfo.remote != via.UdpAddr {
|
||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
||||
hostinfo.logger(f.l).WithField("newAddr", via.UdpAddr).Debug("lighthouse.remote_allow_list denied roaming")
|
||||
return
|
||||
}
|
||||
|
||||
if !hostinfo.lastRoam.IsZero() && via.UdpAddr == hostinfo.lastRoamRemote && time.Since(hostinfo.lastRoam) < RoamingSuppressSeconds*time.Second {
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(f.l).WithField("udpAddr", hostinfo.remote).WithField("newAddr", via.UdpAddr).
|
||||
Debugf("Suppressing roam back to previous remote for %d seconds", RoamingSuppressSeconds)
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("lighthouse.remote_allow_list denied roaming", "newAddr", via.UdpAddr)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo.logger(f.l).WithField("udpAddr", hostinfo.remote).WithField("newAddr", via.UdpAddr).
|
||||
Info("Host roamed to new udp ip/port.")
|
||||
if !hostinfo.lastRoam.IsZero() && via.UdpAddr == hostinfo.lastRoamRemote && time.Since(hostinfo.lastRoam) < RoamingSuppressSeconds*time.Second {
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("Suppressing roam back to previous remote",
|
||||
"suppressSeconds", RoamingSuppressSeconds,
|
||||
"udpAddr", hostinfo.remote,
|
||||
"newAddr", via.UdpAddr,
|
||||
)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
hostinfo.logger(f.l).Info("Host roamed to new udp ip/port.",
|
||||
"udpAddr", hostinfo.remote,
|
||||
"newAddr", via.UdpAddr,
|
||||
)
|
||||
hostinfo.lastRoam = time.Now()
|
||||
hostinfo.lastRoamRemote = hostinfo.remote
|
||||
hostinfo.SetRemote(via.UdpAddr)
|
||||
@@ -268,23 +294,6 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
||||
|
||||
}
|
||||
|
||||
// handleEncrypted returns true if a packet should be processed, false otherwise
|
||||
func (f *Interface) handleEncrypted(ci *ConnectionState, via ViaSender, h *header.H) bool {
|
||||
// If connectionstate does not exist, send a recv error, if possible, to encourage a fast reconnect
|
||||
if ci == nil {
|
||||
if !via.IsRelayed {
|
||||
f.maybeSendRecvError(via.UdpAddr, h.RemoteIndex)
|
||||
}
|
||||
return false
|
||||
}
|
||||
// If the window check fails, refuse to process the packet, but don't send a recv error
|
||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
var (
|
||||
ErrPacketTooShort = errors.New("packet is too short")
|
||||
ErrUnknownIPVersion = errors.New("packet is an unknown ip version")
|
||||
@@ -295,191 +304,16 @@ var (
|
||||
)
|
||||
|
||||
// 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 {
|
||||
if len(data) < 1 {
|
||||
return ErrPacketTooShort
|
||||
var parsed batch.RxParsed
|
||||
if err := batch.ParsePacket(data, incoming, &parsed); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
parsed.Key.Hydrate(fp)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -491,55 +325,100 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []
|
||||
}
|
||||
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
||||
hostinfo.logger(f.l).WithField("header", h).
|
||||
Debugln("dropping out of window packet")
|
||||
return nil, errors.New("out of window packet")
|
||||
return nil, ErrOutOfWindow
|
||||
}
|
||||
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) bool {
|
||||
var err error
|
||||
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
||||
const (
|
||||
ecnNotECT = 0x00
|
||||
ecnECT1 = 0x01
|
||||
ecnECT0 = 0x02
|
||||
ecnCE = 0x03
|
||||
)
|
||||
|
||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||
// 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 {
|
||||
hostinfo.logger(f.l).WithError(err).Error("Failed to decrypt packet")
|
||||
return false
|
||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||
"error", err,
|
||||
"packet", out,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
err = newPacket(out, true, fwPacket)
|
||||
if err != nil {
|
||||
hostinfo.logger(f.l).WithError(err).WithField("packet", out).
|
||||
Warnf("Error while validating inbound packet")
|
||||
return false
|
||||
}
|
||||
|
||||
if !hostinfo.ConnectionState.window.Update(f.l, messageCounter) {
|
||||
hostinfo.logger(f.l).WithField("fwPacket", fwPacket).
|
||||
Debugln("dropping out of window packet")
|
||||
return false
|
||||
}
|
||||
|
||||
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 {
|
||||
// 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
|
||||
f.rejectOutside(out, hostinfo.ConnectionState, hostinfo, nb, packet, q)
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
hostinfo.logger(f.l).WithField("fwPacket", fwPacket).
|
||||
WithField("reason", dropReason).
|
||||
Debugln("dropping inbound packet")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
hostinfo.logger(f.l).Debug("dropping inbound packet",
|
||||
"fwPacket", fwPacket,
|
||||
"reason", dropReason,
|
||||
)
|
||||
}
|
||||
return false
|
||||
return
|
||||
}
|
||||
|
||||
f.connectionManager.In(hostinfo)
|
||||
_, err = f.readers[q].Write(out)
|
||||
err = f.batchers[q].CommitInbound(out, parsedRx)
|
||||
if err != nil {
|
||||
f.l.WithError(err).Error("Failed to write to tun")
|
||||
f.l.Error("Failed to write to tun", "error", err)
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
||||
@@ -553,35 +432,41 @@ func (f *Interface) sendRecvError(endpoint netip.AddrPort, index uint32) {
|
||||
|
||||
b := header.Encode(make([]byte, header.Len), header.Version, header.RecvError, 0, index, 0)
|
||||
_ = f.outside.WriteTo(b, endpoint)
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("index", index).
|
||||
WithField("udpAddr", endpoint).
|
||||
Debug("Recv error sent")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Recv error sent",
|
||||
"index", index,
|
||||
"udpAddr", endpoint,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func (f *Interface) handleRecvError(addr netip.AddrPort, h *header.H) {
|
||||
if !f.acceptRecvErrorConfig.ShouldRecvError(addr) {
|
||||
f.l.WithField("index", h.RemoteIndex).
|
||||
WithField("udpAddr", addr).
|
||||
Debug("Recv error received, ignoring")
|
||||
f.l.Debug("Recv error received, ignoring",
|
||||
"index", h.RemoteIndex,
|
||||
"udpAddr", addr,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
if f.l.Level >= logrus.DebugLevel {
|
||||
f.l.WithField("index", h.RemoteIndex).
|
||||
WithField("udpAddr", addr).
|
||||
Debug("Recv error received")
|
||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||
f.l.Debug("Recv error received",
|
||||
"index", h.RemoteIndex,
|
||||
"udpAddr", addr,
|
||||
)
|
||||
}
|
||||
|
||||
hostinfo := f.hostMap.QueryReverseIndex(h.RemoteIndex)
|
||||
if hostinfo == nil {
|
||||
f.l.WithField("remoteIndex", h.RemoteIndex).Debugln("Did not find remote index in main hostmap")
|
||||
f.l.Debug("Did not find remote index in main hostmap", "remoteIndex", h.RemoteIndex)
|
||||
return
|
||||
}
|
||||
|
||||
if hostinfo.remote.IsValid() && hostinfo.remote != addr {
|
||||
f.l.Infoln("Someone spoofing recv_errors? ", addr, hostinfo.remote)
|
||||
f.l.Info("Someone spoofing recv_errors?",
|
||||
"addr", addr,
|
||||
"hostinfoRemote", hostinfo.remote,
|
||||
)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
+20
-19
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/google/gopacket/layers"
|
||||
|
||||
"github.com/slackhq/nebula/firewall"
|
||||
"github.com/slackhq/nebula/overlay/batch"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/ipv4"
|
||||
@@ -21,13 +22,13 @@ func Test_newPacket(t *testing.T) {
|
||||
|
||||
// length fails
|
||||
err := newPacket([]byte{}, true, p)
|
||||
require.ErrorIs(t, err, ErrPacketTooShort)
|
||||
require.ErrorIs(t, err, batch.ErrPacketTooShort)
|
||||
|
||||
err = newPacket([]byte{0x40}, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv4PacketTooShort)
|
||||
require.ErrorIs(t, err, batch.ErrIPv4PacketTooShort)
|
||||
|
||||
err = newPacket([]byte{0x60}, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
||||
|
||||
// length fail with ip options
|
||||
h := ipv4.Header{
|
||||
@@ -40,15 +41,15 @@ func Test_newPacket(t *testing.T) {
|
||||
|
||||
b, _ := h.Marshal()
|
||||
err = newPacket(b, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
||||
require.ErrorIs(t, err, batch.ErrIPv4InvalidHeaderLength)
|
||||
|
||||
// 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)
|
||||
require.ErrorIs(t, err, ErrUnknownIPVersion)
|
||||
require.ErrorIs(t, err, batch.ErrUnknownIPVersion)
|
||||
|
||||
// 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)
|
||||
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
||||
require.ErrorIs(t, err, batch.ErrIPv4InvalidHeaderLength)
|
||||
|
||||
// account for variable ip header length - incoming
|
||||
h = ipv4.Header{
|
||||
@@ -115,7 +116,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
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
|
||||
// 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
|
||||
// the length byte.
|
||||
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
|
||||
// next layer, missing length byte
|
||||
err = newPacket(buffer.Bytes()[:49], true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
||||
err = nil
|
||||
|
||||
// A good ICMP packet
|
||||
@@ -217,7 +218,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
b = buffer.Bytes()
|
||||
b[6] = 255 // 255 is a reserved protocol number
|
||||
err = newPacket(b, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
||||
|
||||
// A good UDP packet
|
||||
ip = layers.IPv6{
|
||||
@@ -264,7 +265,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
|
||||
// Too short UDP packet
|
||||
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
|
||||
b[6] = byte(layers.IPProtocolTCP)
|
||||
@@ -291,7 +292,7 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
|
||||
// Too short TCP packet
|
||||
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
|
||||
ip = layers.IPv6{
|
||||
@@ -336,12 +337,12 @@ func Test_newPacket_v6(t *testing.T) {
|
||||
|
||||
// Ensure buffer bounds checking during processing
|
||||
err = newPacket(b[:41], true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
||||
|
||||
// Invalid AH header
|
||||
b = buffer.Bytes()
|
||||
err = newPacket(b, true, p)
|
||||
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
||||
}
|
||||
|
||||
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||
@@ -448,7 +449,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||
|
||||
// Too short of a fragment packet
|
||||
err = newPacket(secondFrag[:len(secondFrag)-10], false, p)
|
||||
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
||||
}
|
||||
|
||||
func BenchmarkParseV6(b *testing.B) {
|
||||
@@ -529,7 +530,7 @@ func BenchmarkParseV6(b *testing.B) {
|
||||
|
||||
b.Run("Normal", func(b *testing.B) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -537,7 +538,7 @@ func BenchmarkParseV6(b *testing.B) {
|
||||
|
||||
b.Run("FirstFragment", func(b *testing.B) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -545,7 +546,7 @@ func BenchmarkParseV6(b *testing.B) {
|
||||
|
||||
b.Run("SecondFragment", func(b *testing.B) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -590,7 +591,7 @@ func BenchmarkParseV6(b *testing.B) {
|
||||
|
||||
b.Run("200 HopByHop headers", func(b *testing.B) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user