mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 21:16:59 +02:00
Compare commits
17 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| dfe94c6269 | |||
| 1c601d776a | |||
| 17d8ebff93 | |||
| 612d3ef931 | |||
| 8282a629e5 | |||
| c62f27d4b4 | |||
| f5db77f214 | |||
| b9a7d1edf3 | |||
| d1ea33659a | |||
| 8fdd98f639 | |||
| 45bc0fc055 | |||
| 24af30bd78 | |||
| 1d84b81032 | |||
| b155f4b7e1 | |||
| 194d58cd46 | |||
| a476b1fa07 | |||
| 8b02b8128e |
@@ -1,113 +0,0 @@
|
|||||||
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
|
mv build/*.tar.gz release
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v7
|
uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: linux-latest
|
name: linux-latest
|
||||||
path: release
|
path: release
|
||||||
@@ -32,9 +32,6 @@ jobs:
|
|||||||
build-windows:
|
build-windows:
|
||||||
name: Build Windows
|
name: Build Windows
|
||||||
runs-on: windows-latest
|
runs-on: windows-latest
|
||||||
permissions:
|
|
||||||
id-token: write
|
|
||||||
contents: read
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
@@ -57,15 +54,8 @@ jobs:
|
|||||||
mkdir build\dist\windows
|
mkdir build\dist\windows
|
||||||
mv dist\windows\wintun 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
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v7
|
uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: windows-latest
|
name: windows-latest
|
||||||
path: build
|
path: build
|
||||||
@@ -85,7 +75,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Import certificates
|
- name: Import certificates
|
||||||
if: env.HAS_SIGNING_CREDS == 'true'
|
if: env.HAS_SIGNING_CREDS == 'true'
|
||||||
uses: Apple-Actions/import-codesign-certs@v7
|
uses: Apple-Actions/import-codesign-certs@v6
|
||||||
with:
|
with:
|
||||||
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
p12-file-base64: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_P12_BASE64 }}
|
||||||
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
p12-password: ${{ secrets.APPLE_DEVELOPER_CERTIFICATE_PASSWORD }}
|
||||||
@@ -114,7 +104,7 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v7
|
uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: darwin-latest
|
name: darwin-latest
|
||||||
path: ./release/*
|
path: ./release/*
|
||||||
@@ -138,21 +128,21 @@ jobs:
|
|||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: actions/download-artifact@v8
|
uses: actions/download-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: linux-latest
|
name: linux-latest
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: docker/login-action@v4
|
uses: docker/login-action@v3
|
||||||
with:
|
with:
|
||||||
username: ${{ vars.DOCKERHUB_USERNAME }}
|
username: ${{ vars.DOCKERHUB_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: docker/setup-buildx-action@v4
|
uses: docker/setup-buildx-action@v3
|
||||||
|
|
||||||
- name: Build and push images
|
- name: Build and push images
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
@@ -173,7 +163,7 @@ jobs:
|
|||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v8
|
uses: actions/download-artifact@v7
|
||||||
with:
|
with:
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
|
|||||||
@@ -14,18 +14,10 @@ on:
|
|||||||
- 'go.sum'
|
- 'go.sum'
|
||||||
jobs:
|
jobs:
|
||||||
|
|
||||||
smoke-extra-libvirt:
|
smoke-extra:
|
||||||
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||||
name: ${{ matrix.target }}
|
name: Run extra smoke tests
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
strategy:
|
|
||||||
fail-fast: false
|
|
||||||
matrix:
|
|
||||||
target:
|
|
||||||
- freebsd-amd64
|
|
||||||
- openbsd-amd64
|
|
||||||
- netbsd-amd64
|
|
||||||
- linux-amd64-ipv6disable
|
|
||||||
env:
|
env:
|
||||||
VAGRANT_DEFAULT_PROVIDER: libvirt
|
VAGRANT_DEFAULT_PROVIDER: libvirt
|
||||||
steps:
|
steps:
|
||||||
@@ -48,85 +40,28 @@ jobs:
|
|||||||
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
sudo chmod 666 /var/run/libvirt/libvirt-sock
|
||||||
vagrant plugin install vagrant-libvirt
|
vagrant plugin install vagrant-libvirt
|
||||||
|
|
||||||
- name: ${{ matrix.target }}
|
- name: freebsd-amd64
|
||||||
run: make smoke-vagrant/${{ matrix.target }}
|
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.
|
- name: netbsd-amd64
|
||||||
smoke-extra-virtualbox:
|
run: make smoke-vagrant/netbsd-amd64
|
||||||
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:
|
|
||||||
|
|
||||||
- uses: actions/checkout@v6
|
- name: linux-amd64-ipv6disable
|
||||||
|
run: make smoke-vagrant/linux-amd64-ipv6disable
|
||||||
|
|
||||||
- uses: actions/setup-go@v6
|
# linux-386 runs last because it requires disabling KVM to use VirtualBox,
|
||||||
with:
|
# which prevents libvirt (used by the other tests) from working after this point.
|
||||||
go-version: '1.25'
|
- name: install virtualbox for i386 test
|
||||||
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: |
|
run: |
|
||||||
sudo apt-get update && sudo apt-get install -y vagrant virtualbox
|
sudo apt-get install -y virtualbox
|
||||||
sudo rmmod kvm_amd kvm_intel kvm 2>/dev/null || true
|
sudo rmmod kvm_amd kvm_intel kvm 2>/dev/null || true
|
||||||
|
|
||||||
- name: linux-386
|
- name: linux-386
|
||||||
|
env:
|
||||||
|
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
||||||
run: make smoke-vagrant/linux-386
|
run: make smoke-vagrant/linux-386
|
||||||
|
|
||||||
timeout-minutes: 30
|
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
|
|
||||||
|
|||||||
@@ -1,272 +0,0 @@
|
|||||||
#!/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
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,7 +1,7 @@
|
|||||||
# -*- mode: ruby -*-
|
# -*- mode: ruby -*-
|
||||||
# vi: set ft=ruby :
|
# vi: set ft=ruby :
|
||||||
Vagrant.configure("2") do |config|
|
Vagrant.configure("2") do |config|
|
||||||
config.vm.box = "DefinedNet/netbsd10"
|
config.vm.box = "generic/netbsd9"
|
||||||
|
|
||||||
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
config.vm.synced_folder "../build", "/nebula", type: "rsync"
|
||||||
end
|
end
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ jobs:
|
|||||||
- name: Build test mobile
|
- name: Build test mobile
|
||||||
run: make build-test-mobile
|
run: make build-test-mobile
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v7
|
- uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: e2e packet flow linux-latest
|
name: e2e packet flow linux-latest
|
||||||
path: e2e/mermaid/linux-latest
|
path: e2e/mermaid/linux-latest
|
||||||
@@ -125,7 +125,7 @@ jobs:
|
|||||||
- name: End 2 end
|
- name: End 2 end
|
||||||
run: make e2evv
|
run: make e2evv
|
||||||
|
|
||||||
- uses: actions/upload-artifact@v7
|
- uses: actions/upload-artifact@v6
|
||||||
with:
|
with:
|
||||||
name: e2e packet flow ${{ matrix.os }}
|
name: e2e packet flow ${{ matrix.os }}
|
||||||
path: e2e/mermaid/${{ matrix.os }}
|
path: e2e/mermaid/${{ matrix.os }}
|
||||||
|
|||||||
@@ -2,42 +2,24 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"math"
|
|
||||||
mathbits "math/bits"
|
|
||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
)
|
)
|
||||||
|
|
||||||
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 {
|
type Bits struct {
|
||||||
length uint64
|
length uint64
|
||||||
lengthMask uint64
|
|
||||||
current uint64
|
current uint64
|
||||||
bits []uint64
|
bits []bool
|
||||||
lostCounter metrics.Counter
|
lostCounter metrics.Counter
|
||||||
dupeCounter metrics.Counter
|
dupeCounter metrics.Counter
|
||||||
outOfWindowCounter metrics.Counter
|
outOfWindowCounter metrics.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewBits(length uint64) *Bits {
|
func NewBits(bits 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{
|
b := &Bits{
|
||||||
length: length,
|
length: bits,
|
||||||
lengthMask: length - 1,
|
bits: make([]bool, bits, bits),
|
||||||
bits: make([]uint64, nWords),
|
|
||||||
current: 0,
|
current: 0,
|
||||||
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
lostCounter: metrics.GetOrRegisterCounter("network.packets.lost", nil),
|
||||||
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
dupeCounter: metrics.GetOrRegisterCounter("network.packets.duplicate", nil),
|
||||||
@@ -45,194 +27,71 @@ func NewBits(length uint64) *Bits {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
// There is no counter value 0, mark it to avoid counting a lost packet later.
|
||||||
b.bits[0] = 1
|
b.bits[0] = true
|
||||||
|
b.current = 0
|
||||||
return b
|
return b
|
||||||
}
|
}
|
||||||
|
|
||||||
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 {
|
func (b *Bits) Check(l *slog.Logger, i uint64) bool {
|
||||||
// If i is the next number, return true.
|
// If i is the next number, return true.
|
||||||
if i > b.current {
|
if i > b.current {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
if b.strictlyWithinWindow(i) {
|
// If i is within the window, check if it's been set already.
|
||||||
return !b.get(i)
|
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||||
|
return !b.bits[i%b.length]
|
||||||
}
|
}
|
||||||
|
|
||||||
// Not within the window
|
// Not within the window
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
l.Debug("rejected a packet (top)", "current", b.current, "incoming", i)
|
l.Debug("rejected a packet (top)",
|
||||||
|
"current", b.current,
|
||||||
|
"incoming", i,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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 {
|
func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
||||||
// Fast path: i is the next expected counter. Split out so the function
|
// If i is the next number, return true and update current.
|
||||||
// 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 {
|
if i == b.current+1 {
|
||||||
pos := i & b.lengthMask
|
// Check if the oldest bit was lost since we are shifting the window by 1 and occupying it with this counter
|
||||||
word := pos >> 6
|
// The very first window can only be tracked as lost once we are on the 2nd window or greater
|
||||||
mask := uint64(1) << (pos & 63)
|
if b.bits[i%b.length] == false && i > b.length {
|
||||||
w := b.bits[word]
|
|
||||||
if i > b.length && w&mask == 0 {
|
|
||||||
b.lostCounter.Inc(1)
|
b.lostCounter.Inc(1)
|
||||||
}
|
}
|
||||||
b.bits[word] = w | mask
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
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 is a jump, adjust the window, record lost, update current, and return true
|
||||||
if i > b.current {
|
if i > b.current {
|
||||||
end := i
|
lost := int64(0)
|
||||||
if end > b.current+b.length {
|
// Zero out the bits between the current and the new counter value, limited by the window size,
|
||||||
end = b.current + b.length
|
// since the window is shifting
|
||||||
}
|
for n := b.current + 1; n <= min(i, b.current+b.length); n++ {
|
||||||
count := end - b.current
|
if b.bits[n%b.length] == false && n > b.length {
|
||||||
startPos := (b.current + 1) & b.lengthMask
|
lost++
|
||||||
|
|
||||||
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.clearRange(startPos, count)
|
b.bits[n%b.length] = false
|
||||||
}
|
}
|
||||||
|
|
||||||
// Anything past the new window can never be backfilled, so it's lost.
|
// Only record any skipped packets as a result of the window moving further than the window length
|
||||||
if i > b.current+b.length {
|
// Any loss within the new window will be accounted for in future calls
|
||||||
lost += int64(i - b.current - b.length)
|
lost += max(0, int64(i-b.current-b.length))
|
||||||
}
|
|
||||||
b.lostCounter.Inc(lost)
|
b.lostCounter.Inc(lost)
|
||||||
|
|
||||||
b.set(i)
|
b.bits[i%b.length] = true
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is within the current window but below the current counter, check to see if it's a duplicate
|
// If i is within the current window but below the current counter,
|
||||||
if b.strictlyWithinWindow(i) {
|
// Check to see if it's a duplicate
|
||||||
pos := i & b.lengthMask
|
if i > b.current-b.length || i < b.length && b.current < b.length {
|
||||||
word := pos >> 6
|
if b.current == i || b.bits[i%b.length] == true {
|
||||||
mask := uint64(1) << (pos & 63)
|
|
||||||
w := b.bits[word]
|
|
||||||
if b.current == i || w&mask != 0 {
|
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
l.Debug("Receive window",
|
l.Debug("Receive window",
|
||||||
"accepted", false,
|
"accepted", false,
|
||||||
@@ -245,7 +104,7 @@ func (b *Bits) updateSlow(l *slog.Logger, i uint64) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
b.bits[word] = w | mask
|
b.bits[i%b.length] = true
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+129
-276
@@ -7,79 +7,61 @@ import (
|
|||||||
"github.com/stretchr/testify/assert"
|
"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) {
|
func TestBits(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
assert.EqualValues(t, 16, b.length)
|
|
||||||
|
// make sure it is the right size
|
||||||
|
assert.Len(t, b.bits, 10)
|
||||||
|
|
||||||
// This is initialized to zero - receive one. This should work.
|
// This is initialized to zero - receive one. This should work.
|
||||||
assert.True(t, b.Check(l, 1))
|
assert.True(t, b.Check(l, 1))
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
assert.EqualValues(t, 1, b.current)
|
assert.EqualValues(t, 1, b.current)
|
||||||
g := []bool{true, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
g := []bool{true, true, false, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Receive two
|
// Receive two
|
||||||
assert.True(t, b.Check(l, 2))
|
assert.True(t, b.Check(l, 2))
|
||||||
assert.True(t, b.Update(l, 2))
|
assert.True(t, b.Update(l, 2))
|
||||||
assert.EqualValues(t, 2, b.current)
|
assert.EqualValues(t, 2, b.current)
|
||||||
g = []bool{true, true, true, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
g = []bool{true, true, true, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Receive two again - it will fail
|
// Receive two again - it will fail
|
||||||
assert.False(t, b.Check(l, 2))
|
assert.False(t, b.Check(l, 2))
|
||||||
assert.False(t, b.Update(l, 2))
|
assert.False(t, b.Update(l, 2))
|
||||||
assert.EqualValues(t, 2, b.current)
|
assert.EqualValues(t, 2, b.current)
|
||||||
|
|
||||||
// Jump ahead to 25, which clears the window and sets slot 25%16 = 9.
|
// Jump ahead to 15, which should clear everything and set the 6th element
|
||||||
assert.True(t, b.Check(l, 25))
|
assert.True(t, b.Check(l, 15))
|
||||||
assert.True(t, b.Update(l, 25))
|
assert.True(t, b.Update(l, 15))
|
||||||
assert.EqualValues(t, 25, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, false, false, false, false, true, false, false, false, false, false, false}
|
g = []bool{false, false, false, false, false, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Mark 24, which is in window (current 25, length 16, window covers [10,25]).
|
// Mark 14, which is allowed because it is in the window
|
||||||
assert.True(t, b.Check(l, 24))
|
assert.True(t, b.Check(l, 14))
|
||||||
assert.True(t, b.Update(l, 24))
|
assert.True(t, b.Update(l, 14))
|
||||||
assert.EqualValues(t, 25, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Mark 5, not allowed because 5 <= current-length (25-16=9).
|
// Mark 5, which is not allowed because it is not in the window
|
||||||
assert.False(t, b.Check(l, 5))
|
assert.False(t, b.Check(l, 5))
|
||||||
assert.False(t, b.Update(l, 5))
|
assert.False(t, b.Update(l, 5))
|
||||||
assert.EqualValues(t, 25, b.current)
|
assert.EqualValues(t, 15, b.current)
|
||||||
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
||||||
assert.Equal(t, g, b.snapshot())
|
assert.Equal(t, g, b.bits)
|
||||||
|
|
||||||
// Make sure we handle wrapping around once to the same slot. With
|
// make sure we handle wrapping around once to the current position
|
||||||
// length=16, packets 1 and 17 share slot 1.
|
b = NewBits(10)
|
||||||
b = NewBits(16)
|
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
assert.True(t, b.Update(l, 17))
|
assert.True(t, b.Update(l, 11))
|
||||||
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}, b.snapshot())
|
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false}, b.bits)
|
||||||
|
|
||||||
// Walk through a few windows in order
|
// Walk through a few windows in order
|
||||||
b = NewBits(16)
|
b = NewBits(10)
|
||||||
for i := uint64(1); i <= 100; i++ {
|
for i := uint64(1); i <= 100; i++ {
|
||||||
assert.True(t, b.Check(l, i), "Error while checking %v", 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)
|
assert.True(t, b.Update(l, i), "Error while updating %v", i)
|
||||||
@@ -90,31 +72,24 @@ func TestBits(t *testing.T) {
|
|||||||
|
|
||||||
func TestBitsLargeJumps(t *testing.T) {
|
func TestBitsLargeJumps(t *testing.T) {
|
||||||
l := test.NewLogger()
|
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()
|
b.lostCounter.Clear()
|
||||||
assert.True(t, b.Update(l, 55))
|
|
||||||
assert.Equal(t, int64(39), b.lostCounter.Count())
|
|
||||||
|
|
||||||
// Update(100): clears 16 slots starting at slot 56%16=8. Only slot 7 (for
|
b = NewBits(10)
|
||||||
// packet 55) was set, so 16 - 1 = 15 evicted slots had unset bits.
|
b.lostCounter.Clear()
|
||||||
// Plus 100 - 55 - 16 = 29 packets fell past the window. Total 44.
|
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.True(t, b.Update(l, 100))
|
assert.Equal(t, int64(45), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(39+44), 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, 100)) // We saw packet 55 and 100 and can still track 90,91,92,93,94,95,96,97,98,99
|
||||||
assert.True(t, b.Update(l, 200))
|
assert.Equal(t, int64(89), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(39+44+99), 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())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsDupeCounter(t *testing.T) {
|
func TestBitsDupeCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
@@ -139,117 +114,120 @@ func TestBitsDupeCounter(t *testing.T) {
|
|||||||
|
|
||||||
func TestBitsOutOfWindowCounter(t *testing.T) {
|
func TestBitsOutOfWindowCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Jump to 20 (warmup branch + 4 past-window packets).
|
|
||||||
assert.True(t, b.Update(l, 20))
|
assert.True(t, b.Update(l, 20))
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
// 9 single-step advances, each evicts a slot whose bit was cleared during
|
assert.True(t, b.Update(l, 21))
|
||||||
// the jump above and whose value was never seen, so each contributes 1
|
assert.True(t, b.Update(l, 22))
|
||||||
// to lostCounter.
|
assert.True(t, b.Update(l, 23))
|
||||||
for n := uint64(21); n <= 29; n++ {
|
assert.True(t, b.Update(l, 24))
|
||||||
assert.True(t, b.Update(l, n))
|
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(0), b.outOfWindowCounter.Count())
|
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.False(t, b.Update(l, 0))
|
||||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
// 4 from the Update(20) jump + 9 from 21..29.
|
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
||||||
assert.Equal(t, int64(13), b.lostCounter.Count())
|
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(1), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsLostCounter(t *testing.T) {
|
func TestBitsLostCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Walk 20..29 like the original, just with a bigger window. Same
|
assert.True(t, b.Update(l, 20))
|
||||||
// reasoning as TestBitsOutOfWindowCounter: 4 past-window from Update(20),
|
assert.True(t, b.Update(l, 21))
|
||||||
// then 9 more from the unit advances.
|
assert.True(t, b.Update(l, 22))
|
||||||
for n := uint64(20); n <= 29; n++ {
|
assert.True(t, b.Update(l, 23))
|
||||||
assert.True(t, b.Update(l, n))
|
assert.True(t, b.Update(l, 24))
|
||||||
}
|
assert.True(t, b.Update(l, 25))
|
||||||
assert.Equal(t, int64(13), b.lostCounter.Count())
|
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
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
|
|
||||||
b = NewBits(16)
|
b = NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
// Update(15) clears the warmup window (no lost), sets slot 15.
|
assert.True(t, b.Update(l, 9))
|
||||||
assert.True(t, b.Update(l, 15))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
// 10 will set 0 index, 0 was already set, no lost packets
|
||||||
// Update(16): slot 0 was already set (NewBits seeded it), and 16 is not
|
assert.True(t, b.Update(l, 10))
|
||||||
// strictly > length, so nothing is recorded as lost.
|
|
||||||
assert.True(t, b.Update(l, 16))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
// 11 will set 1 index, 1 was missed, we should see 1 packet lost
|
||||||
// Update(17): we jumped straight from 0 to 15, so slot 1 was cleared
|
assert.True(t, b.Update(l, 11))
|
||||||
// (and never re-set). 17 > 16 is past warmup, so packet 1 is recorded lost.
|
|
||||||
assert.True(t, b.Update(l, 17))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
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))
|
||||||
|
assert.True(t, b.Update(l, 15))
|
||||||
|
assert.True(t, b.Update(l, 16))
|
||||||
|
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())
|
||||||
|
|
||||||
// Fill in 18..30 in single steps. Each i evicts slot i%16. Slots 2..14
|
// Jump ahead by a window size
|
||||||
// were all cleared during Update(15), and we never re-set any of them,
|
assert.True(t, b.Update(l, 29))
|
||||||
// so each i in 18..30 is a fresh lost packet — 13 more.
|
assert.Equal(t, int64(8), b.lostCounter.Count())
|
||||||
for n := uint64(18); n <= 30; n++ {
|
// Now lets walk ahead normally through the window, the missed packets should fill in
|
||||||
assert.True(t, b.Update(l, n))
|
assert.True(t, b.Update(l, 30))
|
||||||
}
|
assert.True(t, b.Update(l, 31))
|
||||||
assert.Equal(t, int64(14), b.lostCounter.Count())
|
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())
|
||||||
|
|
||||||
// Jump ahead by exactly one window size.
|
// Jump ahead by 2 windows, should have recording 1 full window missing
|
||||||
assert.True(t, b.Update(l, 46))
|
assert.True(t, b.Update(l, 58))
|
||||||
// end = min(46, 30+16) = 46, count = 16, all slots cleared. Before the
|
assert.Equal(t, int64(27), b.lostCounter.Count())
|
||||||
// jump every slot 0..15 had been set (Update(15), (16), (17), 18..30),
|
// Now lets walk ahead normally through the window, the missed packets should fill in from this window
|
||||||
// so wasSet=16 and 46 == current+length means no past-window slack:
|
assert.True(t, b.Update(l, 59))
|
||||||
// lost contribution = 0.
|
assert.True(t, b.Update(l, 60))
|
||||||
assert.Equal(t, int64(14), b.lostCounter.Count())
|
assert.True(t, b.Update(l, 61))
|
||||||
|
assert.True(t, b.Update(l, 62))
|
||||||
// Walk 47..55. The Update(46) jump cleared every slot, so only slot 14
|
assert.True(t, b.Update(l, 63))
|
||||||
// (for packet 46) is set when we start. Each subsequent unit step lands
|
assert.True(t, b.Update(l, 64))
|
||||||
// on a slot that was cleared and is past warmup, so it counts as lost.
|
assert.True(t, b.Update(l, 65))
|
||||||
// 9 more = 23.
|
assert.True(t, b.Update(l, 66))
|
||||||
for n := uint64(47); n <= 55; n++ {
|
assert.True(t, b.Update(l, 67))
|
||||||
assert.True(t, b.Update(l, n))
|
// 68 packets tracked, 32 seen, 36 missed
|
||||||
}
|
assert.Equal(t, int64(36), b.lostCounter.Count())
|
||||||
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.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestBitsLostCounterIssue1(t *testing.T) {
|
func TestBitsLostCounterIssue1(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(16)
|
b := NewBits(10)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.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.True(t, b.Update(l, 4))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
@@ -266,7 +244,7 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
|||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 7))
|
assert.True(t, b.Update(l, 7))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
// Skip packet 8.
|
// assert.True(t, b.Update(l, 8))
|
||||||
assert.True(t, b.Update(l, 10))
|
assert.True(t, b.Update(l, 10))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 11))
|
assert.True(t, b.Update(l, 11))
|
||||||
@@ -274,23 +252,9 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
|||||||
|
|
||||||
assert.True(t, b.Update(l, 14))
|
assert.True(t, b.Update(l, 14))
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
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
|
||||||
// Jump to 25. With length=16, slot 25%16=9 corresponds to packet 9
|
assert.True(t, b.Update(l, 19))
|
||||||
// (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())
|
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.True(t, b.Update(l, 12))
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 13))
|
assert.True(t, b.Update(l, 13))
|
||||||
@@ -299,140 +263,29 @@ func TestBitsLostCounterIssue1(t *testing.T) {
|
|||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 16))
|
assert.True(t, b.Update(l, 16))
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
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 and that loss is still recorded once, never
|
// We missed packet 8 above
|
||||||
// double-counted, never zeroed.
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
assert.Equal(t, int64(0), b.outOfWindowCounter.Count())
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestBitsWarmupOvershoot exercises the jump path's warmup arm with an
|
func BenchmarkBits(b *testing.B) {
|
||||||
// overshoot past one full window. NewBits leaves current=0 with only slot 0
|
z := NewBits(10)
|
||||||
// "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++ {
|
for n := 0; n < b.N; n++ {
|
||||||
z.Update(l, uint64(n)+1)
|
for i := range z.bits {
|
||||||
}
|
z.bits[i] = true
|
||||||
}
|
}
|
||||||
|
for i := range z.bits {
|
||||||
|
z.bits[i] = false
|
||||||
|
}
|
||||||
|
|
||||||
// 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,10 +217,6 @@ func (ncp *CAPool) verify(c Certificate, now time.Time, certFp string, signerFp
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if signer.Certificate.Curve() != c.Curve() {
|
|
||||||
return nil, ErrCurveMismatch
|
|
||||||
}
|
|
||||||
|
|
||||||
if signer.Certificate.Expired(now) {
|
if signer.Certificate.Expired(now) {
|
||||||
return nil, ErrRootExpired
|
return nil, ErrRootExpired
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -654,31 +654,3 @@ func TestCertificateV2_Verify_Subnets(t *testing.T) {
|
|||||||
_, err = caPool.VerifyCertificate(time.Now(), c)
|
_, err = caPool.VerifyCertificate(time.Now(), c)
|
||||||
require.NoError(t, err)
|
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,9 +112,6 @@ func (c *certificateV1) CheckSignature(key []byte) bool {
|
|||||||
}
|
}
|
||||||
switch c.details.curve {
|
switch c.details.curve {
|
||||||
case Curve_CURVE25519:
|
case Curve_CURVE25519:
|
||||||
if len(key) != ed25519.PublicKeySize {
|
|
||||||
return false //avoids a panic internal to ed25519
|
|
||||||
}
|
|
||||||
return ed25519.Verify(key, b, c.signature)
|
return ed25519.Verify(key, b, c.signature)
|
||||||
case Curve_P256:
|
case Curve_P256:
|
||||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
||||||
|
|||||||
@@ -151,9 +151,6 @@ func (c *certificateV2) CheckSignature(key []byte) bool {
|
|||||||
|
|
||||||
switch c.curve {
|
switch c.curve {
|
||||||
case Curve_CURVE25519:
|
case Curve_CURVE25519:
|
||||||
if len(key) != ed25519.PublicKeySize {
|
|
||||||
return false //avoids a panic internal to ed25519
|
|
||||||
}
|
|
||||||
return ed25519.Verify(key, b, c.signature)
|
return ed25519.Verify(key, b, c.signature)
|
||||||
case Curve_P256:
|
case Curve_P256:
|
||||||
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
pubKey, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), key)
|
||||||
|
|||||||
@@ -22,7 +22,6 @@ var (
|
|||||||
ErrCaNotFound = errors.New("could not find ca for the certificate")
|
ErrCaNotFound = errors.New("could not find ca for the certificate")
|
||||||
ErrUnknownVersion = errors.New("certificate version unrecognized")
|
ErrUnknownVersion = errors.New("certificate version unrecognized")
|
||||||
ErrCertPubkeyPresent = errors.New("certificate has unexpected pubkey present")
|
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")
|
ErrInvalidPEMBlock = errors.New("input did not contain a valid PEM encoded block")
|
||||||
ErrInvalidPEMCertificateBanner = errors.New("bytes did not contain a proper certificate banner")
|
ErrInvalidPEMCertificateBanner = errors.New("bytes did not contain a proper certificate banner")
|
||||||
|
|||||||
@@ -163,55 +163,3 @@ func P256Keypair() ([]byte, []byte) {
|
|||||||
pubkey := privkey.PublicKey()
|
pubkey := privkey.PublicKey()
|
||||||
return pubkey.Bytes(), privkey.Bytes()
|
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
|
|
||||||
}
|
|
||||||
|
|||||||
+42
-12
@@ -11,6 +11,7 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
@@ -44,16 +45,19 @@ type connectionManager struct {
|
|||||||
inactivityTimeout atomic.Int64
|
inactivityTimeout atomic.Int64
|
||||||
dropInactive atomic.Bool
|
dropInactive atomic.Bool
|
||||||
|
|
||||||
|
metricsTxPunchy metrics.Counter
|
||||||
|
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
func newConnectionManagerFromConfig(l *slog.Logger, c *config.C, hm *HostMap, p *Punchy) *connectionManager {
|
||||||
cm := &connectionManager{
|
cm := &connectionManager{
|
||||||
hostMap: hm,
|
hostMap: hm,
|
||||||
l: l,
|
l: l,
|
||||||
punchy: p,
|
punchy: p,
|
||||||
relayUsed: make(map[uint32]struct{}),
|
relayUsed: make(map[uint32]struct{}),
|
||||||
relayUsedLock: &sync.RWMutex{},
|
relayUsedLock: &sync.RWMutex{},
|
||||||
|
metricsTxPunchy: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
||||||
}
|
}
|
||||||
|
|
||||||
cm.reload(c, true)
|
cm.reload(c, true)
|
||||||
@@ -365,7 +369,7 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
if !outTraffic {
|
if !outTraffic {
|
||||||
// Send a punch packet to keep the NAT state alive
|
// Send a punch packet to keep the NAT state alive
|
||||||
cm.punchy.SendPunch(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
return decision, hostinfo, primary
|
return decision, hostinfo, primary
|
||||||
@@ -396,16 +400,17 @@ func (cm *connectionManager) makeTrafficDecision(localIndex uint32, now time.Tim
|
|||||||
|
|
||||||
// If we aren't sending or receiving traffic then its an unused tunnel and we don't to test the tunnel.
|
// 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.
|
// Just maintain NAT state if configured to do so.
|
||||||
cm.punchy.SendPunch(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// We aren't receiving traffic but we are sending it. The outbound
|
if cm.punchy.GetTargetEverything() {
|
||||||
// traffic itself refreshes the primary remote's NAT state; this
|
// This is similar to the old punchy behavior with a slight optimization.
|
||||||
// fans out to non-primary remotes, but only if target_all_remotes
|
// We aren't receiving traffic but we are sending it, punch on all known
|
||||||
// is configured.
|
// ips in case we need to re-prime NAT state
|
||||||
cm.punchy.SendPunchToAll(hostinfo)
|
cm.sendPunch(hostinfo)
|
||||||
|
}
|
||||||
|
|
||||||
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
if cm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(cm.l).Debug("Tunnel status",
|
hostinfo.logger(cm.l).Debug("Tunnel status",
|
||||||
@@ -507,6 +512,31 @@ 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) {
|
func (cm *connectionManager) tryRehandshake(hostinfo *HostInfo) {
|
||||||
cs := cm.intf.pki.getCertState()
|
cs := cm.intf.pki.getCertState()
|
||||||
curCrt := hostinfo.ConnectionState.myCert
|
curCrt := hostinfo.ConnectionState.myCert
|
||||||
|
|||||||
+15
-10
@@ -7,6 +7,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/overlay/overlaytest"
|
"github.com/slackhq/nebula/overlay/overlaytest"
|
||||||
@@ -46,7 +47,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -64,7 +65,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -79,6 +80,7 @@ func Test_NewConnectionManagerTest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -128,7 +130,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -146,7 +148,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -161,6 +163,7 @@ func Test_NewConnectionManagerTest2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -212,7 +215,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -233,7 +236,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
conf.Settings["tunnels"] = map[string]any{
|
conf.Settings["tunnels"] = map[string]any{
|
||||||
"drop_inactive": true,
|
"drop_inactive": true,
|
||||||
}
|
}
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
assert.True(t, nc.dropInactive.Load())
|
assert.True(t, nc.dropInactive.Load())
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
@@ -246,6 +249,7 @@ func Test_NewConnectionManager_DisconnectInactive(t *testing.T) {
|
|||||||
}
|
}
|
||||||
hostinfo.ConnectionState = &ConnectionState{
|
hostinfo.ConnectionState = &ConnectionState{
|
||||||
myCert: &dummyCert{version: cert.Version1},
|
myCert: &dummyCert{version: cert.Version1},
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|
||||||
@@ -336,9 +340,9 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
cachedPeerCert, err := ncp.VerifyCertificate(now.Add(time.Second), peerCert)
|
cachedPeerCert, err := ncp.VerifyCertificate(now.Add(time.Second), peerCert)
|
||||||
|
|
||||||
cs := &CertState{
|
cs := &CertState{
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{},
|
v1Cert: &dummyCert{},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
lh := newTestLighthouse()
|
lh := newTestLighthouse()
|
||||||
@@ -358,7 +362,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
|
|
||||||
// Create manager
|
// Create manager
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
ifce.connectionManager = nc
|
ifce.connectionManager = nc
|
||||||
@@ -368,6 +372,7 @@ func Test_NewConnectionManagerTest_DisconnectInvalid(t *testing.T) {
|
|||||||
ConnectionState: &ConnectionState{
|
ConnectionState: &ConnectionState{
|
||||||
myCert: &dummyCert{},
|
myCert: &dummyCert{},
|
||||||
peerCert: cachedPeerCert,
|
peerCert: cachedPeerCert,
|
||||||
|
H: &noise.HandshakeState{},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
nc.hostMap.unlockedAddHostInfo(hostinfo, ifce)
|
||||||
|
|||||||
+53
-19
@@ -1,20 +1,23 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/rand"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 8192
|
const ReplayWindow = 1024 //todo I've started seeing out-of-window messages in testing?
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey noiseutil.CipherState
|
eKey *NebulaCipherState
|
||||||
dKey noiseutil.CipherState
|
dKey *NebulaCipherState
|
||||||
|
H *noise.HandshakeState
|
||||||
myCert cert.Certificate
|
myCert cert.Certificate
|
||||||
peerCert *cert.CachedCertificate
|
peerCert *cert.CachedCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
@@ -23,24 +26,55 @@ type ConnectionState struct {
|
|||||||
writeLock sync.Mutex
|
writeLock sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
// newConnectionStateFromResult builds a fully-populated ConnectionState from a
|
func NewConnectionState(cs *CertState, crt cert.Certificate, initiator bool, pattern noise.HandshakePattern) (*ConnectionState, error) {
|
||||||
// completed handshake.Result. It seeds messageCounter and the replay window so
|
var dhFunc noise.DHFunc
|
||||||
// that the post-handshake message indices already used on the wire don't count
|
switch crt.Curve() {
|
||||||
// as missed traffic in the data plane.
|
case cert.Curve_CURVE25519:
|
||||||
func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
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.
|
||||||
ci := &ConnectionState{
|
ci := &ConnectionState{
|
||||||
myCert: r.MyCert,
|
H: hs,
|
||||||
initiator: r.Initiator,
|
initiator: initiator,
|
||||||
peerCert: r.RemoteCert,
|
|
||||||
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
|
||||||
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
|
||||||
window: NewBits(ReplayWindow),
|
window: NewBits(ReplayWindow),
|
||||||
|
myCert: crt,
|
||||||
}
|
}
|
||||||
ci.messageCounter.Add(r.MessageIndex)
|
// always start the counter from 2, as packet 1 and packet 2 are handshake packets.
|
||||||
for i := uint64(1); i <= r.MessageIndex; i++ {
|
ci.messageCounter.Add(2)
|
||||||
ci.window.Update(nil, i)
|
|
||||||
}
|
return ci, nil
|
||||||
return ci
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
func (cs *ConnectionState) MarshalJSON() ([]byte, error) {
|
||||||
|
|||||||
@@ -1,114 +0,0 @@
|
|||||||
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())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
+60
-12
@@ -5,6 +5,8 @@ package nebula
|
|||||||
import (
|
import (
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/google/gopacket"
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
@@ -20,9 +22,7 @@ func (c *Control) WaitForType(msgType header.MessageType, subType header.Message
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
match := h.Type == msgType && h.Subtype == subType
|
if h.Type == msgType && h.Subtype == subType {
|
||||||
p.Release()
|
|
||||||
if match {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -38,9 +38,7 @@ func (c *Control) WaitForTypeByIndex(toIndex uint32, msgType header.MessageType,
|
|||||||
panic(err)
|
panic(err)
|
||||||
}
|
}
|
||||||
pipeTo.InjectUDPPacket(p)
|
pipeTo.InjectUDPPacket(p)
|
||||||
match := h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType
|
if h.RemoteIndex == toIndex && h.Type == msgType && h.Subtype == subType {
|
||||||
p.Release()
|
|
||||||
if match {
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -92,15 +90,65 @@ func (c *Control) GetTunTxChan() <-chan []byte {
|
|||||||
return c.f.inside.(*overlay.TestTun).TxPackets
|
return c.f.inside.(*overlay.TestTun).TxPackets
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectUDPPacket injects a packet into the udp side. We copy internally so the caller keeps ownership of p.
|
// InjectUDPPacket will inject a packet into the udp side of nebula
|
||||||
// The copy comes from the freelist so steady-state alloc is zero.
|
|
||||||
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
func (c *Control) InjectUDPPacket(p *udp.Packet) {
|
||||||
c.f.outside.(*udp.TesterConn).Send(p.Copy())
|
c.f.outside.(*udp.TesterConn).Send(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
// InjectTunPacket pushes an IP packet onto the tun interface.
|
// InjectTunUDPPacket puts a udp packet on the tun interface. Using UDP here because it's a simpler protocol
|
||||||
func (c *Control) InjectTunPacket(packet []byte) {
|
func (c *Control) InjectTunUDPPacket(toAddr netip.Addr, toPort uint16, fromAddr netip.Addr, fromPort uint16, data []byte) {
|
||||||
c.f.inside.(*overlay.TestTun).Send(packet)
|
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())
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Control) GetVpnAddrs() []netip.Addr {
|
func (c *Control) GetVpnAddrs() []netip.Addr {
|
||||||
|
|||||||
@@ -28,7 +28,6 @@ func makeHandshakePacket(from, to netip.AddrPort, subtype header.MessageSubType,
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify the responder correctly handles receiving the same msg1 multiple times
|
// Verify the responder correctly handles receiving the same msg1 multiple times
|
||||||
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
// (retransmission). The duplicate goes through CheckAndComplete -> ErrAlreadySeen
|
||||||
// and the cached response is resent.
|
// and the cached response is resent.
|
||||||
@@ -47,7 +46,7 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me to them")
|
t.Log("Trigger handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
|
|
||||||
t.Log("Grab my msg1")
|
t.Log("Grab my msg1")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -79,7 +78,6 @@ func TestHandshakeRetransmitDuplicate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a truncated handshake packet is ignored and the real
|
// Verify that a truncated handshake packet is ignored and the real
|
||||||
// packet can still complete the handshake.
|
// packet can still complete the handshake.
|
||||||
|
|
||||||
@@ -97,7 +95,7 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake")
|
t.Log("Trigger handshake")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
|
|
||||||
t.Log("Get msg1 and deliver to responder")
|
t.Log("Get msg1 and deliver to responder")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -128,7 +126,6 @@ func TestHandshakeTruncatedPacketRecovery(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// A msg2 arriving with no matching pending index should be silently dropped
|
// A msg2 arriving with no matching pending index should be silently dropped
|
||||||
// with no response sent and no state changes.
|
// with no response sent and no state changes.
|
||||||
|
|
||||||
@@ -146,7 +143,7 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake")
|
t.Log("Complete a normal handshake")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -171,7 +168,6 @@ func TestHandshakeOrphanedMsg2Dropped(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// A handshake packet with an unexpected message counter should be silently
|
// A handshake packet with an unexpected message counter should be silently
|
||||||
// dropped with no side effects and no UDP response.
|
// dropped with no side effects and no UDP response.
|
||||||
|
|
||||||
@@ -203,7 +199,6 @@ func TestHandshakeUnknownMessageCounter(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeUnknownSubtype(t *testing.T) {
|
func TestHandshakeUnknownSubtype(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// A handshake packet with an unknown subtype should be silently dropped.
|
// 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{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -229,7 +224,6 @@ func TestHandshakeUnknownSubtype(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeLateResponse(t *testing.T) {
|
func TestHandshakeLateResponse(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// After a handshake times out, a late response should be silently ignored
|
// After a handshake times out, a late response should be silently ignored
|
||||||
// with no new tunnels created.
|
// with no new tunnels created.
|
||||||
|
|
||||||
@@ -248,7 +242,7 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
|
|
||||||
t.Log("Grab msg1 but don't deliver")
|
t.Log("Grab msg1 but don't deliver")
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
@@ -279,7 +273,6 @@ func TestHandshakeLateResponse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a node rejects a handshake containing its own VPN IP in the
|
// 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.
|
// peer cert. We do this by sending the initiator's own msg1 back to itself.
|
||||||
|
|
||||||
@@ -292,7 +285,7 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
myControl.Start()
|
myControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake from me")
|
t.Log("Trigger handshake from me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(netip.MustParseAddr("10.128.0.2"), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
msg1 := myControl.GetFromUDP(true)
|
msg1 := myControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Drain any handshake retransmits before injecting")
|
t.Log("Drain any handshake retransmits before injecting")
|
||||||
@@ -328,7 +321,6 @@ func TestHandshakeSelfConnectionRejected(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// MessageCounter=0 is not a valid handshake message and should be dropped.
|
// 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{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -349,7 +341,6 @@ func TestHandshakeMessageCounter0Dropped(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRemoteAllowList(t *testing.T) {
|
func TestHandshakeRemoteAllowList(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a handshake from a blocked underlay IP is dropped with no
|
// 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
|
// response and no state changes. Then verify the same packet from an
|
||||||
// allowed IP succeeds.
|
// allowed IP succeeds.
|
||||||
@@ -375,7 +366,7 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Trigger handshake from them")
|
t.Log("Trigger handshake from them")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
msg1 := theirControl.GetFromUDP(true)
|
msg1 := theirControl.GetFromUDP(true)
|
||||||
|
|
||||||
t.Log("Rewrite the source to a blocked IP and inject")
|
t.Log("Rewrite the source to a blocked IP and inject")
|
||||||
@@ -408,7 +399,6 @@ func TestHandshakeRemoteAllowList(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
// When a duplicate msg1 arrives via ErrAlreadySeen, verify the tunnel
|
||||||
// remains functional and hostmap index count is stable.
|
// remains functional and hostmap index count is stable.
|
||||||
|
|
||||||
@@ -426,7 +416,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
defer r.RenderFlow()
|
defer r.RenderFlow()
|
||||||
|
|
||||||
t.Log("Complete a normal handshake via the router")
|
t.Log("Complete a normal handshake via the router")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi"))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
assertTunnel(t, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
||||||
|
|
||||||
@@ -437,7 +427,7 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
originalRemote := hi.CurrentRemote
|
originalRemote := hi.CurrentRemote
|
||||||
|
|
||||||
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
t.Log("Re-trigger traffic to cause a new handshake attempt (ErrAlreadySeen)")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("roam"))
|
||||||
r.RouteForAllUntilTxTun(theirControl)
|
r.RouteForAllUntilTxTun(theirControl)
|
||||||
|
|
||||||
t.Log("Verify tunnel still works")
|
t.Log("Verify tunnel still works")
|
||||||
@@ -455,7 +445,6 @@ func TestHandshakeAlreadySeenPreferredRemote(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that when the wrong host responds, the cached packets are
|
// Verify that when the wrong host responds, the cached packets are
|
||||||
// transferred to the new handshake, the evil tunnel is closed, evil's
|
// transferred to the new handshake, the evil tunnel is closed, evil's
|
||||||
// address is blocked, and the correct tunnel is eventually established.
|
// address is blocked, and the correct tunnel is eventually established.
|
||||||
@@ -475,8 +464,8 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Send multiple packets to them (cached during handshake)")
|
t.Log("Send multiple packets to them (cached during handshake)")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet1"))
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("packet2"))
|
||||||
|
|
||||||
t.Log("Route until evil tunnel is closed")
|
t.Log("Route until evil tunnel is closed")
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
@@ -519,7 +508,6 @@ func TestHandshakeWrongResponderPacketStore(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestHandshakeRelayComplete(t *testing.T) {
|
func TestHandshakeRelayComplete(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
// Verify that a relay handshake completes correctly and relay state is
|
// Verify that a relay handshake completes correctly and relay state is
|
||||||
// properly maintained on all three nodes.
|
// properly maintained on all three nodes.
|
||||||
|
|
||||||
@@ -540,7 +528,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger handshake via relay")
|
t.Log("Trigger handshake via relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi via relay"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi via relay"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -568,7 +556,7 @@ func TestHandshakeRelayComplete(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
// NOTE: Relay V1 cert + IPv6 rejection is not tested here because
|
||||||
// BuildTunUDPPacket from a V4 node to a V6 address panics in the test
|
// InjectTunUDPPacket from a V4 node to a V6 address panics in the test
|
||||||
// framework. The check is in handshake_manager.go handleOutbound relay
|
// framework. The check is in handshake_manager.go handleOutbound relay
|
||||||
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
// logic (lines ~304-313): if the relay host has a V1 cert and either
|
||||||
// address is IPv6, the relay is skipped.
|
// address is IPv6, the relay is skipped.
|
||||||
|
|||||||
+30
-67
@@ -16,7 +16,6 @@ import (
|
|||||||
"github.com/slackhq/nebula/cert_test"
|
"github.com/slackhq/nebula/cert_test"
|
||||||
"github.com/slackhq/nebula/e2e/router"
|
"github.com/slackhq/nebula/e2e/router"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
@@ -40,22 +39,11 @@ func BenchmarkHotPath(b *testing.B) {
|
|||||||
r.CancelFlowLogs()
|
r.CancelFlowLogs()
|
||||||
|
|
||||||
assertTunnel(b, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), myControl, theirControl, r)
|
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()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunPacket(prebuilt)
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
// Release the TUN-side bytes back to the harness freelist; the bench
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
// just confirms a packet arrived, the contents aren't inspected.
|
|
||||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -83,15 +71,11 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
assertTunnel(b, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), theirControl, myControl, r)
|
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()
|
b.ResetTimer()
|
||||||
|
|
||||||
for n := 0; n < b.N; n++ {
|
for n := 0; n < b.N; n++ {
|
||||||
myControl.InjectTunPacket(prebuilt)
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
overlay.ReleaseTunBuf(r.RouteForAllUntilTxTun(theirControl))
|
_ = r.RouteForAllUntilTxTun(theirControl)
|
||||||
}
|
}
|
||||||
|
|
||||||
myControl.Stop()
|
myControl.Stop()
|
||||||
@@ -100,7 +84,6 @@ func BenchmarkHotPathRelay(b *testing.B) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshake(t *testing.T) {
|
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{})
|
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)
|
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)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -113,7 +96,7 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(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")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -151,7 +134,6 @@ func TestGoodHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeNoOverlap(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{})
|
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)
|
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!
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "them", "2001::69/24", nil) //look ma, cross-stack!
|
||||||
@@ -187,7 +169,6 @@ func TestGoodHandshakeNoOverlap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshake(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{})
|
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)
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me", "10.128.0.100/24", nil)
|
||||||
@@ -207,7 +188,7 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -264,7 +245,6 @@ func TestWrongResponderHandshake(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestWrongResponderHandshakeStaticHostMap(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{})
|
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)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.99/24", nil)
|
||||||
@@ -289,7 +269,7 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
evilControl.Start()
|
evilControl.Start()
|
||||||
|
|
||||||
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
t.Log("Start the handshake process, we will route until we see the evil tunnel closed")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
@@ -347,7 +327,6 @@ func TestWrongResponderHandshakeStaticHostMap(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1Race(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
|
// 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
|
// But will eventually collapse down to a single tunnel
|
||||||
|
|
||||||
@@ -368,8 +347,8 @@ func TestStage1Race(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake to start on both me and them")
|
t.Log("Trigger a handshake to start on both me and them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(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")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
t.Log("Get both stage 1 handshake packets")
|
t.Log("Get both stage 1 handshake packets")
|
||||||
myHsForThem := myControl.GetFromUDP(true)
|
myHsForThem := myControl.GetFromUDP(true)
|
||||||
@@ -428,7 +407,6 @@ func TestStage1Race(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceLoser(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{})
|
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)
|
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)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -446,7 +424,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -457,7 +435,7 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
myHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me again"))
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me again"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -478,7 +456,6 @@ func TestUncleanShutdownRaceLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUncleanShutdownRaceWinner(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{})
|
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)
|
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)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", nil)
|
||||||
@@ -496,7 +473,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
r.Log("Trigger a handshake from me to them")
|
r.Log("Trigger a handshake from me to them")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
||||||
@@ -508,7 +485,7 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.Indexes = map[uint32]*nebula.HostInfo{}
|
||||||
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
theirHostmap.RemoteIndexes = map[uint32]*nebula.HostInfo{}
|
||||||
|
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them again"))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them again"), p, theirVpnIpNet[0].Addr(), myVpnIpNet[0].Addr(), 80, 80)
|
||||||
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
r.RenderHostmaps("Derp hostmaps", myControl, theirControl)
|
||||||
@@ -530,7 +507,6 @@ func TestUncleanShutdownRaceWinner(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelays(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{})
|
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}})
|
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}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -551,7 +527,7 @@ func TestRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -560,7 +536,6 @@ func TestRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRelaysDontCareAboutIps(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{})
|
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}})
|
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}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "2001::9999/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -581,7 +556,7 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -590,7 +565,6 @@ func TestRelaysDontCareAboutIps(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestReestablishRelays(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{})
|
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}})
|
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}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -611,14 +585,14 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnIpNet[0].Addr(), theirVpnIpNet[0].Addr(), 80, 80)
|
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")
|
t.Log("Ensure packet traversal from them to me via the relay")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -633,7 +607,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
for curIndexes >= start {
|
for curIndexes >= start {
|
||||||
curIndexes = len(myControl.GetHostmap().Indexes)
|
curIndexes = len(myControl.GetHostmap().Indexes)
|
||||||
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
r.Logf("Wait for the dead index to go away:start=%v indexes, current=%v indexes", start, curIndexes)
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me should fail")))
|
myControl.InjectTunUDPPacket(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 {
|
r.RouteForAllExitFunc(func(p *udp.Packet, c *nebula.Control) router.ExitType {
|
||||||
return router.RouteAndExit
|
return router.RouteAndExit
|
||||||
@@ -650,7 +624,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
myControl.InjectLightHouseAddr(relayVpnIpNet[0].Addr(), relayUdpAddr)
|
||||||
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
myControl.InjectRelays(theirVpnIpNet[0].Addr(), []netip.Addr{relayVpnIpNet[0].Addr()})
|
||||||
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
relayControl.InjectLightHouseAddr(theirVpnIpNet[0].Addr(), theirUdpAddr)
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(theirControl)
|
p = r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -685,7 +659,7 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
t.Log("Assert the tunnel works the other way, too")
|
t.Log("Assert the tunnel works the other way, too")
|
||||||
for {
|
for {
|
||||||
t.Log("RouteForAllUntilTxTun")
|
t.Log("RouteForAllUntilTxTun")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -722,7 +696,6 @@ func TestReestablishRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays(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
|
//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{})
|
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}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -755,8 +728,8 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
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")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(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")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
r.Log("Wait for a packet from them to me")
|
r.Log("Wait for a packet from them to me")
|
||||||
p := r.RouteForAllUntilTxTun(myControl)
|
p := r.RouteForAllUntilTxTun(myControl)
|
||||||
@@ -770,7 +743,6 @@ func TestStage1RaceRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestStage1RaceRelays2(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
|
//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{})
|
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}})
|
myControl, myVpnIpNet, myUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.1/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -803,8 +775,8 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
assertTunnel(t, theirVpnIpNet[0].Addr(), relayVpnIpNet[0].Addr(), theirControl, relayControl, r)
|
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")
|
r.Log("Trigger a handshake from both them and me via relay to them and me")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(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")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(myControl, header.Control, header.MessageNone)
|
||||||
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
//r.RouteUntilAfterMsgType(theirControl, header.Control, header.MessageNone)
|
||||||
@@ -847,7 +819,6 @@ func TestStage1RaceRelays2(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelays(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{})
|
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}})
|
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}})
|
relayControl, relayVpnIpNet, relayUdpAddr, relayConfig := newSimpleServer(cert.Version1, ca, caKey, "relay ", "10.128.0.128/24", m{"relay": m{"am_relay": true}})
|
||||||
@@ -868,7 +839,7 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -951,7 +922,6 @@ func TestRehandshakingRelays(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingRelaysPrimary(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
|
// 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{})
|
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}})
|
myControl, myVpnIpNet, _, _ := newSimpleServer(cert.Version1, ca, caKey, "me ", "10.128.0.128/24", m{"relay": m{"use_relays": true}})
|
||||||
@@ -973,7 +943,7 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
@@ -1056,7 +1026,6 @@ func TestRehandshakingRelaysPrimary(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshaking(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{})
|
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)
|
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)
|
theirControl, theirVpnIpNet, theirUdpAddr, theirConfig := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.1/24", nil)
|
||||||
@@ -1152,7 +1121,6 @@ func TestRehandshaking(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRehandshakingLoser(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
|
// 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
|
// 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{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -1251,7 +1219,6 @@ func TestRehandshakingLoser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRaceRegression(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
|
// 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
|
// 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
|
// caused a cross-linked hostinfo
|
||||||
@@ -1275,8 +1242,8 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
//them rx stage:2 initiatorIndex=120607833 responderIndex=4209862089
|
||||||
|
|
||||||
t.Log("Start both handshakes")
|
t.Log("Start both handshakes")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnIpNet[0].Addr(), 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(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")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, theirVpnIpNet[0].Addr(), 80, []byte("Hi from them"))
|
||||||
|
|
||||||
t.Log("Get both stage 1")
|
t.Log("Get both stage 1")
|
||||||
myStage1ForThem := myControl.GetFromUDP(true)
|
myStage1ForThem := myControl.GetFromUDP(true)
|
||||||
@@ -1312,7 +1279,6 @@ func TestRaceRegression(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithLighthouse(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{})
|
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}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "10.128.0.1/24, ff::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1353,7 +1319,6 @@ func TestV2NonPrimaryWithLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestV2NonPrimaryWithOffNetLighthouse(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{})
|
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}})
|
lhControl, lhVpnIpNet, lhUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "lh ", "2001::1/64", m{"lighthouse": m{"am_lighthouse": true}})
|
||||||
|
|
||||||
@@ -1394,7 +1359,6 @@ func TestV2NonPrimaryWithOffNetLighthouse(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestLighthouseUpdateOnReload(t *testing.T) {
|
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{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version2, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
|
|
||||||
// Create the lighthouse
|
// Create the lighthouse
|
||||||
@@ -1470,7 +1434,6 @@ func TestLighthouseUpdateOnReload(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
||||||
t.Parallel()
|
|
||||||
unsafePrefix := "192.168.6.0/24"
|
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{})
|
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)
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServerWithUdpAndUnsafeNetworks(cert.Version2, ca, caKey, "spooky", "10.128.0.2/24", netip.MustParseAddrPort("10.64.0.2:4242"), unsafePrefix, nil)
|
||||||
@@ -1492,7 +1455,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
t.Log("Send a udp packet through to begin standing up the tunnel, this should come out the other side")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(spookyDest, 80, myVpnIpNet[0].Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
t.Log("Have them consume my stage 0 packet. They have a tunnel now")
|
||||||
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
theirControl.InjectUDPPacket(myControl.GetFromUDP(true))
|
||||||
@@ -1520,7 +1483,7 @@ func TestGoodHandshakeUnsafeDest(t *testing.T) {
|
|||||||
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), myCachedPacket, myVpnIpNet[0].Addr(), spookyDest, 80, 80)
|
||||||
|
|
||||||
//reply
|
//reply
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman")))
|
theirControl.InjectTunUDPPacket(myVpnIpNet[0].Addr(), 80, spookyDest, 80, []byte("Hi from the spookyman"))
|
||||||
//wait for reply
|
//wait for reply
|
||||||
theirControl.WaitForType(1, 0, myControl)
|
theirControl.WaitForType(1, 0, myControl)
|
||||||
theirCachedPacket := myControl.GetFromTun(true)
|
theirCachedPacket := myControl.GetFromTun(true)
|
||||||
|
|||||||
+2
-57
@@ -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) {
|
func assertTunnel(t testing.TB, vpnIpA, vpnIpB netip.Addr, controlA, controlB *nebula.Control, r *router.R) {
|
||||||
// Send a packet from them to me
|
// Send a packet from them to me
|
||||||
controlB.InjectTunPacket(BuildTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B")))
|
controlB.InjectTunUDPPacket(vpnIpA, 80, vpnIpB, 90, []byte("Hi from B"))
|
||||||
bPacket := r.RouteForAllUntilTxTun(controlA)
|
bPacket := r.RouteForAllUntilTxTun(controlA)
|
||||||
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
assertUdpPacket(t, []byte("Hi from B"), bPacket, vpnIpB, vpnIpA, 90, 80)
|
||||||
|
|
||||||
// And once more from me to them
|
// And once more from me to them
|
||||||
controlA.InjectTunPacket(BuildTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A")))
|
controlA.InjectTunUDPPacket(vpnIpB, 80, vpnIpA, 90, []byte("Hello from A"))
|
||||||
aPacket := r.RouteForAllUntilTxTun(controlB)
|
aPacket := r.RouteForAllUntilTxTun(controlB)
|
||||||
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
assertUdpPacket(t, []byte("Hello from A"), aPacket, vpnIpA, vpnIpB, 90, 80)
|
||||||
}
|
}
|
||||||
@@ -408,58 +408,3 @@ func testLogLevelName() string {
|
|||||||
}
|
}
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
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()
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,47 +0,0 @@
|
|||||||
//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)
|
|
||||||
}
|
|
||||||
+54
-188
@@ -13,7 +13,6 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"sort"
|
"sort"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -25,19 +24,6 @@ import (
|
|||||||
"golang.org/x/exp/maps"
|
"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 {
|
type R struct {
|
||||||
// Simple map of the ip:port registered on a control to the control
|
// Simple map of the ip:port registered on a control to the control
|
||||||
// Basically a router, right?
|
// Basically a router, right?
|
||||||
@@ -48,28 +34,12 @@ type R struct {
|
|||||||
|
|
||||||
// A last used map, if an inbound packet hit the inNat map then
|
// 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
|
// all return packets should use the same last used inbound address for the outbound sender
|
||||||
outNat map[outNatKey]netip.AddrPort
|
// map[from address + ":" + to address] => ip:port to rewrite in the udp packet to receiver
|
||||||
|
outNat map[string]netip.AddrPort
|
||||||
|
|
||||||
// A map of vpn ip to the nebula control it belongs to
|
// A map of vpn ip to the nebula control it belongs to
|
||||||
vpnControls map[netip.Addr]*nebula.Control
|
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
|
ignoreFlows []ignoreFlow
|
||||||
flow []flowEntry
|
flow []flowEntry
|
||||||
|
|
||||||
@@ -149,7 +119,7 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
controls: make(map[netip.AddrPort]*nebula.Control),
|
controls: make(map[netip.AddrPort]*nebula.Control),
|
||||||
vpnControls: make(map[netip.Addr]*nebula.Control),
|
vpnControls: make(map[netip.Addr]*nebula.Control),
|
||||||
inNat: make(map[netip.AddrPort]*nebula.Control),
|
inNat: make(map[netip.AddrPort]*nebula.Control),
|
||||||
outNat: make(map[outNatKey]netip.AddrPort),
|
outNat: make(map[string]netip.AddrPort),
|
||||||
flow: []flowEntry{},
|
flow: []flowEntry{},
|
||||||
ignoreFlows: []ignoreFlow{},
|
ignoreFlows: []ignoreFlow{},
|
||||||
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
fn: filepath.Join("mermaid", fmt.Sprintf("%s.md", t.Name())),
|
||||||
@@ -183,10 +153,8 @@ func NewR(t testing.TB, controls ...*nebula.Control) *R {
|
|||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-clockSource.C:
|
case <-clockSource.C:
|
||||||
r.Lock()
|
|
||||||
r.renderHostmaps("clock tick")
|
r.renderHostmaps("clock tick")
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
r.Unlock()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
@@ -212,21 +180,15 @@ 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.
|
// RenderFlow renders the packet flow seen up until now and stops further automatic renders from happening.
|
||||||
func (r *R) RenderFlow() {
|
func (r *R) RenderFlow() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
r.Lock()
|
|
||||||
defer r.Unlock()
|
|
||||||
r.renderFlow()
|
r.renderFlow()
|
||||||
}
|
}
|
||||||
|
|
||||||
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
// CancelFlowLogs stops flow logs from being tracked and destroys any logs already collected
|
||||||
func (r *R) CancelFlowLogs() {
|
func (r *R) CancelFlowLogs() {
|
||||||
r.cancelRender()
|
r.cancelRender()
|
||||||
r.Lock()
|
|
||||||
r.flow = nil
|
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() {
|
func (r *R) renderFlow() {
|
||||||
if r.flow == nil {
|
if r.flow == nil {
|
||||||
return
|
return
|
||||||
@@ -472,157 +434,68 @@ func (r *R) RouteUntilTxTun(sender *nebula.Control, receiver *nebula.Control) []
|
|||||||
panic("No control for udp tx " + a.String())
|
panic("No control for udp tx " + a.String())
|
||||||
}
|
}
|
||||||
fp := r.unlockedInjectFlow(sender, c, p, false)
|
fp := r.unlockedInjectFlow(sender, c, p, false)
|
||||||
c.InjectUDPPacket(p) // copies internally; original is ours to release
|
c.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on the receiver's tun.
|
// RouteForAllUntilTxTun will route for everyone and return when a packet is seen on receivers tun
|
||||||
// If a control's UDP TX address can't be matched to a registered control, we panic.
|
// If the router doesn't have the nebula controller for that address, 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 {
|
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)
|
sc := make([]reflect.SelectCase, len(r.controls)+1)
|
||||||
cm := make([]*nebula.Control, len(r.controls)+1)
|
cm := make([]*nebula.Control, len(r.controls)+1)
|
||||||
sc[0] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(receiver.GetTunTxChan())}
|
|
||||||
cm[0] = receiver
|
i := 0
|
||||||
i := 1
|
sc[i] = reflect.SelectCase{
|
||||||
|
Dir: reflect.SelectRecv,
|
||||||
|
Chan: reflect.ValueOf(receiver.GetTunTxChan()),
|
||||||
|
Send: reflect.Value{},
|
||||||
|
}
|
||||||
|
cm[i] = receiver
|
||||||
|
|
||||||
|
i++
|
||||||
for _, c := range r.controls {
|
for _, c := range r.controls {
|
||||||
sc[i] = reflect.SelectCase{Dir: reflect.SelectRecv, Chan: reflect.ValueOf(c.GetUDPTxChan())}
|
sc[i] = reflect.SelectCase{
|
||||||
|
Dir: reflect.SelectRecv,
|
||||||
|
Chan: reflect.ValueOf(c.GetUDPTxChan()),
|
||||||
|
Send: reflect.Value{},
|
||||||
|
}
|
||||||
|
|
||||||
cm[i] = c
|
cm[i] = c
|
||||||
i++
|
i++
|
||||||
}
|
}
|
||||||
r.selRecvCtl = receiver
|
|
||||||
r.selCases = sc
|
for {
|
||||||
r.selCtls = cm
|
x, rx, _ := reflect.Select(sc)
|
||||||
return sc, cm
|
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()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
// RouteExitFunc will call the whatDo func with each udp packet from sender.
|
||||||
@@ -649,7 +522,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -657,7 +529,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -670,7 +541,6 @@ func (r *R) RouteExitFunc(sender *nebula.Control, whatDo ExitFunc) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -771,7 +641,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
switch e {
|
switch e {
|
||||||
case ExitNow:
|
case ExitNow:
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case RouteAndExit:
|
case RouteAndExit:
|
||||||
@@ -779,7 +648,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
fp.WasReceived()
|
fp.WasReceived()
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
return
|
return
|
||||||
|
|
||||||
case KeepRouting:
|
case KeepRouting:
|
||||||
@@ -791,7 +659,6 @@ func (r *R) RouteForAllExitFunc(whatDo ExitFunc) {
|
|||||||
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
panic(fmt.Sprintf("Unknown exitFunc return: %v", e))
|
||||||
}
|
}
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -835,20 +702,19 @@ func (r *R) FlushAll() {
|
|||||||
}
|
}
|
||||||
receiver.InjectUDPPacket(p)
|
receiver.InjectUDPPacket(p)
|
||||||
r.Unlock()
|
r.Unlock()
|
||||||
p.Release()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// getControl performs or seeds NAT translation and returns the control for toAddr, p from fields may change
|
// 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
|
// 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 {
|
func (r *R) getControl(fromAddr, toAddr netip.AddrPort, p *udp.Packet) *nebula.Control {
|
||||||
if newAddr, ok := r.outNat[outNatKey{from: fromAddr, to: toAddr}]; ok {
|
if newAddr, ok := r.outNat[fromAddr.String()+":"+toAddr.String()]; ok {
|
||||||
p.From = newAddr
|
p.From = newAddr
|
||||||
}
|
}
|
||||||
|
|
||||||
c, ok := r.inNat[toAddr]
|
c, ok := r.inNat[toAddr]
|
||||||
if ok {
|
if ok {
|
||||||
r.outNat[outNatKey{from: c.GetUDPAddr(), to: fromAddr}] = toAddr
|
r.outNat[c.GetUDPAddr().String()+":"+fromAddr.String()] = toAddr
|
||||||
return c
|
return c
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,125 +0,0 @@
|
|||||||
//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)
|
|
||||||
}
|
|
||||||
+2
-8
@@ -19,7 +19,6 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestDropInactiveTunnels(t *testing.T) {
|
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
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -64,7 +63,6 @@ func TestDropInactiveTunnels(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertUpgrade(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
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -159,7 +157,6 @@ func TestCertUpgrade(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertDowngrade(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
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -258,7 +255,6 @@ func TestCertDowngrade(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCertMismatchCorrection(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
|
// The goal of this test is to ensure the shortest inactivity timeout will close the tunnel on both sides
|
||||||
// under ideal conditions
|
// under ideal conditions
|
||||||
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
ca, _, caKey, _ := cert_test.NewTestCaCert(cert.Version1, cert.Curve_CURVE25519, time.Now(), time.Now().Add(10*time.Minute), nil, nil, []string{})
|
||||||
@@ -326,7 +322,6 @@ func TestCertMismatchCorrection(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCrossStackRelaysWork(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{})
|
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}})
|
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}})
|
relayControl, relayVpnIpNet, relayUdpAddr, _ := newSimpleServer(cert.Version2, ca, caKey, "relay ", "10.128.0.128/24,fc00::128/64", m{"relay": m{"am_relay": true}})
|
||||||
@@ -355,14 +350,14 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
theirControl.Start()
|
theirControl.Start()
|
||||||
|
|
||||||
t.Log("Trigger a handshake from me to them via the relay")
|
t.Log("Trigger a handshake from me to them via the relay")
|
||||||
myControl.InjectTunPacket(BuildTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me")))
|
myControl.InjectTunUDPPacket(theirVpnV6.Addr(), 80, myVpnV6.Addr(), 80, []byte("Hi from me"))
|
||||||
|
|
||||||
p := r.RouteForAllUntilTxTun(theirControl)
|
p := r.RouteForAllUntilTxTun(theirControl)
|
||||||
r.Log("Assert the tunnel works")
|
r.Log("Assert the tunnel works")
|
||||||
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from me"), p, myVpnV6.Addr(), theirVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
t.Log("reply?")
|
t.Log("reply?")
|
||||||
theirControl.InjectTunPacket(BuildTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them")))
|
theirControl.InjectTunUDPPacket(myVpnV6.Addr(), 80, theirVpnV6.Addr(), 80, []byte("Hi from them"))
|
||||||
p = r.RouteForAllUntilTxTun(myControl)
|
p = r.RouteForAllUntilTxTun(myControl)
|
||||||
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
assertUdpPacket(t, []byte("Hi from them"), p, theirVpnV6.Addr(), myVpnV6.Addr(), 80, 80)
|
||||||
|
|
||||||
@@ -374,7 +369,6 @@ func TestCrossStackRelaysWork(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestCloseTunnelAuthenticated(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{})
|
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"}})
|
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"}})
|
theirControl, theirVpnIpNet, theirUdpAddr, _ := newSimpleServer(cert.Version1, ca, caKey, "them", "10.128.0.2/24", m{"tunnels": m{"drop_inactive": true, "inactivity_timeout": "10m"}})
|
||||||
|
|||||||
@@ -1,125 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -138,14 +138,6 @@ listen:
|
|||||||
# max, net.core.rmem_max and net.core.wmem_max
|
# max, net.core.rmem_max and net.core.wmem_max
|
||||||
#read_buffer: 10485760
|
#read_buffer: 10485760
|
||||||
#write_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
|
# 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
|
# 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.
|
# 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.
|
||||||
@@ -171,21 +163,17 @@ listen:
|
|||||||
|
|
||||||
punchy:
|
punchy:
|
||||||
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
# Continues to punch inbound/outbound at a regular interval to avoid expiration of firewall nat mappings
|
||||||
# This setting is reloadable.
|
|
||||||
punch: true
|
punch: true
|
||||||
|
|
||||||
# respond means that a node you are trying to reach will connect back out to you if your hole punching fails
|
# 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
|
# this is extremely useful if one node is behind a difficult nat, such as a symmetric NAT
|
||||||
# Default is false
|
# Default is false
|
||||||
# This setting is reloadable.
|
|
||||||
#respond: true
|
#respond: true
|
||||||
|
|
||||||
# delays a punch response for misbehaving NATs, default is 1 second.
|
# delays a punch response for misbehaving NATs, default is 1 second.
|
||||||
# This setting is reloadable.
|
|
||||||
#delay: 1s
|
#delay: 1s
|
||||||
|
|
||||||
# set the delay before attempting punchy.respond. Default is 5 seconds. respond must be true to take effect.
|
# 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
|
#respond_delay: 5s
|
||||||
|
|
||||||
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
# Cipher allows you to choose between the available ciphers for your network. Options are chachapoly or aes
|
||||||
@@ -294,24 +282,6 @@ tun:
|
|||||||
# metric: 100
|
# metric: 100
|
||||||
# install: true
|
# 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
|
# 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.
|
# in nebula configuration files. Default false, not reloadable.
|
||||||
#use_system_route_table: false
|
#use_system_route_table: false
|
||||||
|
|||||||
+25
-44
@@ -80,8 +80,8 @@ type firewallMetrics struct {
|
|||||||
type FirewallConntrack struct {
|
type FirewallConntrack struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
|
|
||||||
Conns map[firewall.PacketKey]*conn
|
Conns map[firewall.Packet]*conn
|
||||||
TimerWheel *TimerWheel[firewall.PacketKey]
|
TimerWheel *TimerWheel[firewall.Packet]
|
||||||
}
|
}
|
||||||
|
|
||||||
// FirewallTable is the entry point for a rule, the evaluation order is:
|
// FirewallTable is the entry point for a rule, the evaluation order is:
|
||||||
@@ -166,8 +166,8 @@ func NewFirewall(l *slog.Logger, tcpTimeout, UDPTimeout, defaultTimeout time.Dur
|
|||||||
|
|
||||||
return &Firewall{
|
return &Firewall{
|
||||||
Conntrack: &FirewallConntrack{
|
Conntrack: &FirewallConntrack{
|
||||||
Conns: make(map[firewall.PacketKey]*conn),
|
Conns: make(map[firewall.Packet]*conn),
|
||||||
TimerWheel: NewTimerWheel[firewall.PacketKey](tmin, tmax),
|
TimerWheel: NewTimerWheel[firewall.Packet](tmin, tmax),
|
||||||
},
|
},
|
||||||
InRules: newFirewallTable(),
|
InRules: newFirewallTable(),
|
||||||
OutRules: newFirewallTable(),
|
OutRules: newFirewallTable(),
|
||||||
@@ -422,27 +422,12 @@ var ErrNoMatchingRule = errors.New("no matching rule in firewall table")
|
|||||||
|
|
||||||
// Drop returns an error if the packet should be dropped, explaining why. It
|
// Drop returns an error if the packet should be dropped, explaining why. It
|
||||||
// returns nil if the packet should not be dropped.
|
// returns nil if the packet should not be dropped.
|
||||||
//
|
func (f *Firewall) Drop(fp firewall.Packet, incoming bool, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) error {
|
||||||
// key is the dense conntrack key — used as-is for the inConns fast path
|
// Check if we spoke to this tuple, if we did then allow this packet
|
||||||
// without touching fp at all. fp is the rich Packet form rule matching
|
if f.inConns(fp, h, caPool, localCache) {
|
||||||
// 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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Conntrack miss → rule matching needs the rich Packet form. Hydrate
|
|
||||||
// from the key if the caller passed a zero-valued fp (the inbound path
|
|
||||||
// after batch.ParsePacket). Outbound callers Hydrate themselves and
|
|
||||||
// skip this hop.
|
|
||||||
if !fp.LocalAddr.IsValid() {
|
|
||||||
key.Hydrate(fp)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Make sure remote address matches nebula certificate, and determine how to treat it
|
// Make sure remote address matches nebula certificate, and determine how to treat it
|
||||||
if h.networks == nil {
|
if h.networks == nil {
|
||||||
// Simple case: Certificate has one address and no unsafe networks
|
// Simple case: Certificate has one address and no unsafe networks
|
||||||
@@ -482,13 +467,13 @@ func (f *Firewall) Drop(key firewall.PacketKey, fp *firewall.Packet, incoming bo
|
|||||||
}
|
}
|
||||||
|
|
||||||
// We now know which firewall table to check against
|
// We now know which firewall table to check against
|
||||||
if !table.match(*fp, incoming, h.ConnectionState.peerCert, caPool) {
|
if !table.match(fp, incoming, h.ConnectionState.peerCert, caPool) {
|
||||||
f.metrics(incoming).droppedNoRule.Inc(1)
|
f.metrics(incoming).droppedNoRule.Inc(1)
|
||||||
return ErrNoMatchingRule
|
return ErrNoMatchingRule
|
||||||
}
|
}
|
||||||
|
|
||||||
// We always want to conntrack since it is a faster operation
|
// We always want to conntrack since it is a faster operation
|
||||||
f.addConn(key, fp.Protocol, incoming)
|
f.addConn(fp, incoming)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -517,9 +502,9 @@ func (f *Firewall) EmitStats() {
|
|||||||
metrics.GetOrRegisterGauge("firewall.rules.hash", nil).Update(int64(f.GetRuleHashFNV()))
|
metrics.GetOrRegisterGauge("firewall.rules.hash", nil).Update(int64(f.GetRuleHashFNV()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) bool {
|
func (f *Firewall) inConns(fp firewall.Packet, h *HostInfo, caPool *cert.CAPool, localCache firewall.ConntrackCache) bool {
|
||||||
if localCache != nil {
|
if localCache != nil {
|
||||||
if _, ok := localCache[key]; ok {
|
if _, ok := localCache[fp]; ok {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -532,7 +517,7 @@ func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAP
|
|||||||
f.evict(ep)
|
f.evict(ep)
|
||||||
}
|
}
|
||||||
|
|
||||||
c, ok := conntrack.Conns[key]
|
c, ok := conntrack.Conns[fp]
|
||||||
|
|
||||||
if !ok {
|
if !ok {
|
||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
@@ -541,11 +526,7 @@ func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAP
|
|||||||
|
|
||||||
if c.rulesVersion != f.rulesVersion {
|
if c.rulesVersion != f.rulesVersion {
|
||||||
// This conntrack entry was for an older rule set, validate
|
// This conntrack entry was for an older rule set, validate
|
||||||
// it still passes with the current rule set. Rule matching needs
|
// it still passes with the current rule set
|
||||||
// the rich Packet form, so hydrate from key.
|
|
||||||
var fp firewall.Packet
|
|
||||||
key.Hydrate(&fp)
|
|
||||||
|
|
||||||
table := f.OutRules
|
table := f.OutRules
|
||||||
if c.incoming {
|
if c.incoming {
|
||||||
table = f.InRules
|
table = f.InRules
|
||||||
@@ -561,7 +542,7 @@ func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAP
|
|||||||
"oldRulesVersion", c.rulesVersion,
|
"oldRulesVersion", c.rulesVersion,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
delete(conntrack.Conns, key)
|
delete(conntrack.Conns, fp)
|
||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -578,7 +559,7 @@ func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAP
|
|||||||
c.rulesVersion = f.rulesVersion
|
c.rulesVersion = f.rulesVersion
|
||||||
}
|
}
|
||||||
|
|
||||||
switch key.Protocol {
|
switch fp.Protocol {
|
||||||
case firewall.ProtoTCP:
|
case firewall.ProtoTCP:
|
||||||
c.Expires = time.Now().Add(f.TCPTimeout)
|
c.Expires = time.Now().Add(f.TCPTimeout)
|
||||||
case firewall.ProtoUDP:
|
case firewall.ProtoUDP:
|
||||||
@@ -590,17 +571,17 @@ func (f *Firewall) inConns(key firewall.PacketKey, h *HostInfo, caPool *cert.CAP
|
|||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
|
|
||||||
if localCache != nil {
|
if localCache != nil {
|
||||||
localCache[key] = struct{}{}
|
localCache[fp] = struct{}{}
|
||||||
}
|
}
|
||||||
|
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Firewall) addConn(key firewall.PacketKey, protocol uint8, incoming bool) {
|
func (f *Firewall) addConn(fp firewall.Packet, incoming bool) {
|
||||||
var timeout time.Duration
|
var timeout time.Duration
|
||||||
c := &conn{}
|
c := &conn{}
|
||||||
|
|
||||||
switch protocol {
|
switch fp.Protocol {
|
||||||
case firewall.ProtoTCP:
|
case firewall.ProtoTCP:
|
||||||
timeout = f.TCPTimeout
|
timeout = f.TCPTimeout
|
||||||
case firewall.ProtoUDP:
|
case firewall.ProtoUDP:
|
||||||
@@ -611,9 +592,9 @@ func (f *Firewall) addConn(key firewall.PacketKey, protocol uint8, incoming bool
|
|||||||
|
|
||||||
conntrack := f.Conntrack
|
conntrack := f.Conntrack
|
||||||
conntrack.Lock()
|
conntrack.Lock()
|
||||||
if _, ok := conntrack.Conns[key]; !ok {
|
if _, ok := conntrack.Conns[fp]; !ok {
|
||||||
conntrack.TimerWheel.Advance(time.Now())
|
conntrack.TimerWheel.Advance(time.Now())
|
||||||
conntrack.TimerWheel.Add(key, timeout)
|
conntrack.TimerWheel.Add(fp, timeout)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Record which rulesVersion allowed this connection, so we can retest after
|
// Record which rulesVersion allowed this connection, so we can retest after
|
||||||
@@ -621,16 +602,16 @@ func (f *Firewall) addConn(key firewall.PacketKey, protocol uint8, incoming bool
|
|||||||
c.incoming = incoming
|
c.incoming = incoming
|
||||||
c.rulesVersion = f.rulesVersion
|
c.rulesVersion = f.rulesVersion
|
||||||
c.Expires = time.Now().Add(timeout)
|
c.Expires = time.Now().Add(timeout)
|
||||||
conntrack.Conns[key] = c
|
conntrack.Conns[fp] = c
|
||||||
conntrack.Unlock()
|
conntrack.Unlock()
|
||||||
}
|
}
|
||||||
|
|
||||||
// Evict checks if a conntrack entry has expired, if so it is removed, if not it is re-added to the wheel
|
// Evict checks if a conntrack entry has expired, if so it is removed, if not it is re-added to the wheel
|
||||||
// Caller must own the connMutex lock!
|
// Caller must own the connMutex lock!
|
||||||
func (f *Firewall) evict(key firewall.PacketKey) {
|
func (f *Firewall) evict(p firewall.Packet) {
|
||||||
// Are we still tracking this conn?
|
// Are we still tracking this conn?
|
||||||
conntrack := f.Conntrack
|
conntrack := f.Conntrack
|
||||||
t, ok := conntrack.Conns[key]
|
t, ok := conntrack.Conns[p]
|
||||||
if !ok {
|
if !ok {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
@@ -640,12 +621,12 @@ func (f *Firewall) evict(key firewall.PacketKey) {
|
|||||||
// Timeout is in the future, re-add the timer
|
// Timeout is in the future, re-add the timer
|
||||||
if newT > 0 {
|
if newT > 0 {
|
||||||
conntrack.TimerWheel.Advance(time.Now())
|
conntrack.TimerWheel.Advance(time.Now())
|
||||||
conntrack.TimerWheel.Add(key, newT)
|
conntrack.TimerWheel.Add(p, newT)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// This conn is done
|
// This conn is done
|
||||||
delete(conntrack.Conns, key)
|
delete(conntrack.Conns, p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool {
|
func (ft *FirewallTable) match(p firewall.Packet, incoming bool, c *cert.CachedCertificate, caPool *cert.CAPool) bool {
|
||||||
|
|||||||
+4
-8
@@ -5,15 +5,11 @@ import (
|
|||||||
"log/slog"
|
"log/slog"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// ConntrackCache is used as a local routine cache to know if a given flow
|
// ConntrackCache is used as a local routine cache to know if a given flow
|
||||||
// has been seen in the conntrack table. Keyed on PacketKey (dense form)
|
// has been seen in the conntrack table.
|
||||||
// rather than Packet so the lookup hashes raw bytes instead of the
|
type ConntrackCache map[Packet]struct{}
|
||||||
// unique.Handle each netip.Addr in Packet carries.
|
|
||||||
type ConntrackCache map[PacketKey]struct{}
|
|
||||||
|
|
||||||
type ConntrackCacheTicker struct {
|
type ConntrackCacheTicker struct {
|
||||||
cacheV uint64
|
cacheV uint64
|
||||||
@@ -60,8 +56,8 @@ func (c *ConntrackCacheTicker) Get() ConntrackCache {
|
|||||||
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
if tick := c.cacheTick.Load(); tick != c.cacheV {
|
||||||
c.cacheV = tick
|
c.cacheV = tick
|
||||||
if ll := len(c.cache); ll > 0 {
|
if ll := len(c.cache); ll > 0 {
|
||||||
if c.l.Enabled(context.Background(), logging.LevelTrace) {
|
if c.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
c.l.Log(context.Background(), logging.LevelTrace, "resetting conntrack cache", "len", ll)
|
c.l.Debug("resetting conntrack cache", "len", ll)
|
||||||
}
|
}
|
||||||
c.cache = make(ConntrackCache, ll)
|
c.cache = make(ConntrackCache, ll)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/logging"
|
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
)
|
)
|
||||||
@@ -23,7 +22,7 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
|||||||
cache: make(ConntrackCache, cacheLen),
|
cache: make(ConntrackCache, cacheLen),
|
||||||
}
|
}
|
||||||
for i := 0; i < cacheLen; i++ {
|
for i := 0; i < cacheLen; i++ {
|
||||||
c.cache[PacketKey{LocalPort: uint16(i) + 1}] = struct{}{}
|
c.cache[Packet{LocalPort: uint16(i) + 1}] = struct{}{}
|
||||||
}
|
}
|
||||||
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
|
c.cacheTick.Store(1) // cacheV starts at 0, so Get() takes the reset path
|
||||||
return c
|
return c
|
||||||
@@ -31,27 +30,27 @@ func newFixedTicker(t *testing.T, l *slog.Logger, cacheLen int) *ConntrackCacheT
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_TextFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 3)
|
c := newFixedTicker(t, l, 3)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.Equal(t, "level=DEBUG-4 msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
assert.Equal(t, "level=DEBUG msg=\"resetting conntrack cache\" len=3\n", buf.String())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
func TestConntrackCacheTicker_Get_JSONFormat(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewJSONLoggerWithOutput(buf, logging.LevelTrace)
|
l := test.NewJSONLoggerWithOutput(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 2)
|
c := newFixedTicker(t, l, 2)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|
||||||
assert.JSONEq(t, `{"level":"DEBUG-4","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
assert.JSONEq(t, `{"level":"DEBUG","msg":"resetting conntrack cache","len":2}`, strings.TrimSpace(buf.String()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietBelowDebug(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelInfo)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 5)
|
c := newFixedTicker(t, l, 5)
|
||||||
c.Get()
|
c.Get()
|
||||||
@@ -61,7 +60,7 @@ func TestConntrackCacheTicker_Get_QuietBelowTrace(t *testing.T) {
|
|||||||
|
|
||||||
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
func TestConntrackCacheTicker_Get_QuietWhenCacheEmpty(t *testing.T) {
|
||||||
buf := &bytes.Buffer{}
|
buf := &bytes.Buffer{}
|
||||||
l := test.NewLoggerWithOutputAndLevel(buf, logging.LevelTrace)
|
l := test.NewLoggerWithOutputAndLevel(buf, slog.LevelDebug)
|
||||||
|
|
||||||
c := newFixedTicker(t, l, 0)
|
c := newFixedTicker(t, l, 0)
|
||||||
c.Get()
|
c.Get()
|
||||||
|
|||||||
@@ -19,25 +19,6 @@ const (
|
|||||||
PortFragment = -1 // Special value for matching `port: fragment`
|
PortFragment = -1 // Special value for matching `port: fragment`
|
||||||
)
|
)
|
||||||
|
|
||||||
// PacketKey is the firewall's conntrack and ConntrackCache map key — the
|
|
||||||
// dense form of the 5-tuple plus the protocol and fragment flag the
|
|
||||||
// firewall actually discriminates flows on. Kept separate from Packet so
|
|
||||||
// the conntrack-hit fast path doesn't pay for hashing the unique.Handle
|
|
||||||
// each netip.Addr carries, and so the inbound parser can skip the
|
|
||||||
// AddrFrom4/AddrFrom16 calls until rule matching actually needs them.
|
|
||||||
//
|
|
||||||
// Superset of the coalescer's flowKey shape (same 5-tuple, just in
|
|
||||||
// Local/Remote orientation rather than wire src/dst).
|
|
||||||
type PacketKey struct {
|
|
||||||
LocalAddr [16]byte
|
|
||||||
RemoteAddr [16]byte
|
|
||||||
LocalPort uint16
|
|
||||||
RemotePort uint16
|
|
||||||
IsV6 bool
|
|
||||||
Protocol uint8
|
|
||||||
Fragment bool
|
|
||||||
}
|
|
||||||
|
|
||||||
type Packet struct {
|
type Packet struct {
|
||||||
LocalAddr netip.Addr
|
LocalAddr netip.Addr
|
||||||
RemoteAddr netip.Addr
|
RemoteAddr netip.Addr
|
||||||
@@ -50,61 +31,6 @@ type Packet struct {
|
|||||||
Fragment bool
|
Fragment bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// Key derives a PacketKey from a populated Packet. Used by the few code
|
|
||||||
// paths that have a Packet but no Key in hand (e.g. tests). Both inbound
|
|
||||||
// and outbound production parsers write straight into a PacketKey via
|
|
||||||
// batch.ParsePacket, so this function is rarely on the hot path.
|
|
||||||
func (fp *Packet) Key() PacketKey {
|
|
||||||
k := PacketKey{
|
|
||||||
Protocol: fp.Protocol,
|
|
||||||
Fragment: fp.Fragment,
|
|
||||||
}
|
|
||||||
k.LocalPort = fp.LocalPort
|
|
||||||
k.RemotePort = fp.RemotePort
|
|
||||||
k.IsV6 = !fp.LocalAddr.Is4()
|
|
||||||
if k.IsV6 {
|
|
||||||
k.LocalAddr = fp.LocalAddr.As16()
|
|
||||||
k.RemoteAddr = fp.RemoteAddr.As16()
|
|
||||||
} else {
|
|
||||||
v4 := fp.LocalAddr.As4()
|
|
||||||
copy(k.LocalAddr[:4], v4[:])
|
|
||||||
v4 = fp.RemoteAddr.As4()
|
|
||||||
copy(k.RemoteAddr[:4], v4[:])
|
|
||||||
}
|
|
||||||
return k
|
|
||||||
}
|
|
||||||
|
|
||||||
// Hydrate fills fp's netip.Addr fields and copies the rest from k. Called
|
|
||||||
// by the firewall slow path when conntrack misses and rule matching needs
|
|
||||||
// the rich Packet form (CIDR lookups, family checks). The fast path skips
|
|
||||||
// this entirely.
|
|
||||||
func (k *PacketKey) Hydrate(fp *Packet) {
|
|
||||||
fp.LocalPort = k.LocalPort
|
|
||||||
fp.RemotePort = k.RemotePort
|
|
||||||
fp.Protocol = k.Protocol
|
|
||||||
fp.Fragment = k.Fragment
|
|
||||||
if k.IsV6 {
|
|
||||||
fp.LocalAddr = netip.AddrFrom16(k.LocalAddr)
|
|
||||||
fp.RemoteAddr = netip.AddrFrom16(k.RemoteAddr)
|
|
||||||
} else {
|
|
||||||
var v4 [4]byte
|
|
||||||
copy(v4[:], k.LocalAddr[:4])
|
|
||||||
fp.LocalAddr = netip.AddrFrom4(v4)
|
|
||||||
copy(v4[:], k.RemoteAddr[:4])
|
|
||||||
fp.RemoteAddr = netip.AddrFrom4(v4)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (k *PacketKey) GetRemoteAddr() netip.Addr {
|
|
||||||
if k.IsV6 {
|
|
||||||
return netip.AddrFrom16(k.RemoteAddr)
|
|
||||||
} else {
|
|
||||||
var v4 [4]byte
|
|
||||||
copy(v4[:], k.RemoteAddr[:4])
|
|
||||||
return netip.AddrFrom4(v4)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (fp *Packet) Copy() *Packet {
|
func (fp *Packet) Copy() *Packet {
|
||||||
return &Packet{
|
return &Packet{
|
||||||
LocalAddr: fp.LocalAddr,
|
LocalAddr: fp.LocalAddr,
|
||||||
|
|||||||
+54
-54
@@ -211,44 +211,44 @@ func TestFirewall_Drop(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil))
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
// Allow outbound because conntrack
|
// Allow outbound because conntrack
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
||||||
|
|
||||||
// test remote mismatch
|
// test remote mismatch
|
||||||
oldRemote := p.RemoteAddr
|
oldRemote := p.RemoteAddr
|
||||||
p.RemoteAddr = netip.MustParseAddr("1.2.3.10")
|
p.RemoteAddr = netip.MustParseAddr("1.2.3.10")
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP)
|
||||||
p.RemoteAddr = oldRemote
|
p.RemoteAddr = oldRemote
|
||||||
|
|
||||||
// ensure signer doesn't get in the way of group checks
|
// ensure signer doesn't get in the way of group checks
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caSha doesn't drop on match
|
// test caSha doesn't drop on match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
|
|
||||||
// ensure ca name doesn't get in the way of group checks
|
// ensure ca name doesn't get in the way of group checks
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caName doesn't drop on match
|
// test caName doesn't drop on match
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_DropV6(t *testing.T) {
|
func TestFirewall_DropV6(t *testing.T) {
|
||||||
@@ -289,44 +289,44 @@ func TestFirewall_DropV6(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
assert.Equal(t, ErrNoMatchingRule, fw.Drop(p, false, &h, cp, nil))
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
// Allow outbound because conntrack
|
// Allow outbound because conntrack
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
||||||
|
|
||||||
// test remote mismatch
|
// test remote mismatch
|
||||||
oldRemote := p.RemoteAddr
|
oldRemote := p.RemoteAddr
|
||||||
p.RemoteAddr = netip.MustParseAddr("fd12::56")
|
p.RemoteAddr = netip.MustParseAddr("fd12::56")
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrInvalidRemoteIP)
|
||||||
p.RemoteAddr = oldRemote
|
p.RemoteAddr = oldRemote
|
||||||
|
|
||||||
// ensure signer doesn't get in the way of group checks
|
// ensure signer doesn't get in the way of group checks
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caSha doesn't drop on match
|
// test caSha doesn't drop on match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "", "signer-shasum-bad"))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "", "signer-shasum"))
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
|
|
||||||
// ensure ca name doesn't get in the way of group checks
|
// ensure ca name doesn't get in the way of group checks
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good-bad", ""))
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// test caName doesn't drop on match
|
// test caName doesn't drop on match
|
||||||
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
cp.CAs["signer-shasum"] = &cert.CachedCertificate{Certificate: &dummyCert{name: "ca-good"}}
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, &c)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"nope"}, "", "", "", "ca-good-bad", ""))
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 0, 0, []string{"default-group"}, "", "", "", "ca-good", ""))
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkFirewallTable_match(b *testing.B) {
|
func BenchmarkFirewallTable_match(b *testing.B) {
|
||||||
@@ -533,10 +533,10 @@ func TestFirewall_Drop2(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// h1/c1 lacks the proper groups
|
// h1/c1 lacks the proper groups
|
||||||
require.ErrorIs(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil), ErrNoMatchingRule)
|
require.ErrorIs(t, fw.Drop(p, true, &h1, cp, nil), ErrNoMatchingRule)
|
||||||
// c has the proper groups
|
// c has the proper groups
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_Drop3(t *testing.T) {
|
func TestFirewall_Drop3(t *testing.T) {
|
||||||
@@ -613,18 +613,18 @@ func TestFirewall_Drop3(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// c1 should pass because host match
|
// c1 should pass because host match
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h1, cp, nil))
|
||||||
// c2 should pass because ca sha match
|
// c2 should pass because ca sha match
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h2, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h2, cp, nil))
|
||||||
// c3 should fail because no match
|
// c3 should fail because no match
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, true, &h3, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, true, &h3, cp, nil), ErrNoMatchingRule)
|
||||||
|
|
||||||
// Test a remote address match
|
// Test a remote address match
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "1.2.3.4/24", "", "", ""))
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h1, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_Drop3V6(t *testing.T) {
|
func TestFirewall_Drop3V6(t *testing.T) {
|
||||||
@@ -661,7 +661,7 @@ func TestFirewall_Drop3V6(t *testing.T) {
|
|||||||
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw := NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", ""))
|
require.NoError(t, fw.AddRule(true, firewall.ProtoAny, 1, 1, []string{}, "", "fd12::34/120", "", "", ""))
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_DropConntrackReload(t *testing.T) {
|
func TestFirewall_DropConntrackReload(t *testing.T) {
|
||||||
@@ -702,12 +702,12 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
|
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, true, &h, cp, nil))
|
||||||
// Allow outbound because conntrack
|
// Allow outbound because conntrack
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
||||||
|
|
||||||
oldFw := fw
|
oldFw := fw
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
@@ -716,7 +716,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
fw.rulesVersion = oldFw.rulesVersion + 1
|
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
|
||||||
// Allow outbound because conntrack and new rules allow port 10
|
// Allow outbound because conntrack and new rules allow port 10
|
||||||
require.NoError(t, fw.Drop(p.Key(), &p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(p, false, &h, cp, nil))
|
||||||
|
|
||||||
oldFw = fw
|
oldFw = fw
|
||||||
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
fw = NewFirewall(l, time.Second, time.Minute, time.Hour, c.Certificate)
|
||||||
@@ -725,7 +725,7 @@ func TestFirewall_DropConntrackReload(t *testing.T) {
|
|||||||
fw.rulesVersion = oldFw.rulesVersion + 1
|
fw.rulesVersion = oldFw.rulesVersion + 1
|
||||||
|
|
||||||
// Drop outbound because conntrack doesn't match new ruleset
|
// Drop outbound because conntrack doesn't match new ruleset
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
||||||
@@ -770,12 +770,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero ports", func(t *testing.T) {
|
t.Run("nonzero ports", func(t *testing.T) {
|
||||||
@@ -783,12 +783,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -800,12 +800,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero ports, still blocked", func(t *testing.T) {
|
t.Run("nonzero ports, still blocked", func(t *testing.T) {
|
||||||
@@ -813,12 +813,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero, matching ports, still blocked", func(t *testing.T) {
|
t.Run("nonzero, matching ports, still blocked", func(t *testing.T) {
|
||||||
@@ -826,12 +826,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 80
|
p.LocalPort = 80
|
||||||
p.RemotePort = 80
|
p.RemotePort = 80
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, true, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, true, &h, cp, nil), ErrNoMatchingRule)
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
t.Run("Any proto, any port", func(t *testing.T) {
|
t.Run("Any proto, any port", func(t *testing.T) {
|
||||||
@@ -843,12 +843,12 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0
|
p.LocalPort = 0
|
||||||
p.RemotePort = 0
|
p.RemotePort = 0
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("nonzero ports, allowed", func(t *testing.T) {
|
t.Run("nonzero ports, allowed", func(t *testing.T) {
|
||||||
@@ -857,15 +857,15 @@ func TestFirewall_ICMPPortBehavior(t *testing.T) {
|
|||||||
p.LocalPort = 0xabcd
|
p.LocalPort = 0xabcd
|
||||||
p.RemotePort = 0x1234
|
p.RemotePort = 0x1234
|
||||||
// Drop outbound
|
// Drop outbound
|
||||||
assert.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
assert.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
// Allow inbound
|
// Allow inbound
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, true, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, true, &h, cp, nil))
|
||||||
//now also allow outbound
|
//now also allow outbound
|
||||||
require.NoError(t, fw.Drop(p.Key(), p, false, &h, cp, nil))
|
require.NoError(t, fw.Drop(*p, false, &h, cp, nil))
|
||||||
//different ID is blocked
|
//different ID is blocked
|
||||||
p.RemotePort++
|
p.RemotePort++
|
||||||
require.Equal(t, fw.Drop(p.Key(), p, false, &h, cp, nil), ErrNoMatchingRule)
|
require.Equal(t, fw.Drop(*p, false, &h, cp, nil), ErrNoMatchingRule)
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -913,7 +913,7 @@ func TestFirewall_DropIPSpoofing(t *testing.T) {
|
|||||||
Protocol: firewall.ProtoUDP,
|
Protocol: firewall.ProtoUDP,
|
||||||
Fragment: false,
|
Fragment: false,
|
||||||
}
|
}
|
||||||
assert.Equal(t, fw.Drop(p.Key(), &p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
assert.Equal(t, fw.Drop(p, true, &h1, cp, nil), ErrInvalidRemoteIP)
|
||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkLookup(b *testing.B) {
|
func BenchmarkLookup(b *testing.B) {
|
||||||
@@ -1033,7 +1033,7 @@ func TestNewFirewallFromConfig(t *testing.T) {
|
|||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
// Test a bad rule definition
|
// Test a bad rule definition
|
||||||
c := &dummyCert{}
|
c := &dummyCert{}
|
||||||
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil, "aes")
|
cs, err := newCertState(cert.Version2, nil, c, false, cert.Curve_CURVE25519, nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
conf := config.NewC(test.NewLogger())
|
conf := config.NewC(test.NewLogger())
|
||||||
@@ -1327,7 +1327,7 @@ func (c *testcase) Test(t *testing.T, fw *Firewall) {
|
|||||||
t.Helper()
|
t.Helper()
|
||||||
cp := cert.NewCAPool()
|
cp := cert.NewCAPool()
|
||||||
resetConntrack(fw)
|
resetConntrack(fw)
|
||||||
err := fw.Drop(c.p.Key(), &c.p, true, c.h, cp, nil)
|
err := fw.Drop(c.p, true, c.h, cp, nil)
|
||||||
if c.err == nil {
|
if c.err == nil {
|
||||||
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
require.NoError(t, err, "failed to not drop remote address %s", c.p.RemoteAddr)
|
||||||
} else {
|
} else {
|
||||||
@@ -1519,6 +1519,6 @@ func (mf *mockFirewall) AddRule(incoming bool, proto uint8, startPort int32, end
|
|||||||
|
|
||||||
func resetConntrack(fw *Firewall) {
|
func resetConntrack(fw *Firewall) {
|
||||||
fw.Conntrack.Lock()
|
fw.Conntrack.Lock()
|
||||||
fw.Conntrack.Conns = map[firewall.PacketKey]*conn{}
|
fw.Conntrack.Conns = map[firewall.Packet]*conn{}
|
||||||
fw.Conntrack.Unlock()
|
fw.Conntrack.Unlock()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ require (
|
|||||||
github.com/armon/go-radix v1.0.0
|
github.com/armon/go-radix v1.0.0
|
||||||
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
github.com/cyberdelia/go-metrics-graphite v0.0.0-20161219230853-39f87cc3b432
|
||||||
github.com/flynn/noise v1.1.0
|
github.com/flynn/noise v1.1.0
|
||||||
github.com/gaissmai/bart v0.26.1
|
github.com/gaissmai/bart v0.26.0
|
||||||
github.com/gogo/protobuf v1.3.2
|
github.com/gogo/protobuf v1.3.2
|
||||||
github.com/google/gopacket v1.1.19
|
github.com/google/gopacket v1.1.19
|
||||||
github.com/kardianos/service v1.2.4
|
github.com/kardianos/service v1.2.4
|
||||||
@@ -22,11 +22,10 @@ require (
|
|||||||
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
github.com/stefanberger/go-pkcs11uri v0.0.0-20230803200340-78284954bff6
|
||||||
github.com/stretchr/testify v1.11.1
|
github.com/stretchr/testify v1.11.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
go.uber.org/goleak v1.3.0
|
|
||||||
go.yaml.in/yaml/v3 v3.0.4
|
go.yaml.in/yaml/v3 v3.0.4
|
||||||
golang.org/x/crypto v0.50.0
|
golang.org/x/crypto v0.50.0
|
||||||
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
golang.org/x/exp v0.0.0-20230725093048-515e97ebf090
|
||||||
golang.org/x/net v0.53.0
|
golang.org/x/net v0.52.0
|
||||||
golang.org/x/sync v0.20.0
|
golang.org/x/sync v0.20.0
|
||||||
golang.org/x/sys v0.43.0
|
golang.org/x/sys v0.43.0
|
||||||
golang.org/x/term v0.42.0
|
golang.org/x/term v0.42.0
|
||||||
@@ -43,7 +42,6 @@ require (
|
|||||||
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
github.com/cespare/xxhash/v2 v2.3.0 // indirect
|
||||||
github.com/davecgh/go-spew v1.1.1 // indirect
|
github.com/davecgh/go-spew v1.1.1 // indirect
|
||||||
github.com/google/btree v1.1.2 // indirect
|
github.com/google/btree v1.1.2 // indirect
|
||||||
github.com/guptarohit/asciigraph v0.9.0 // indirect
|
|
||||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect
|
||||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||||
github.com/prometheus/client_model v0.6.2 // indirect
|
github.com/prometheus/client_model v0.6.2 // indirect
|
||||||
|
|||||||
@@ -26,8 +26,8 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
|
|||||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||||
github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
github.com/flynn/noise v1.1.0 h1:KjPQoQCEFdZDiP03phOvGi11+SVVhBG2wOWAorLsstg=
|
||||||
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
github.com/flynn/noise v1.1.0/go.mod h1:xbMo+0i6+IGbYdJhF31t2eR1BIU0CYc12+BNAKwUTag=
|
||||||
github.com/gaissmai/bart v0.26.1 h1:+w4rnLGNlA2GDVn382Tfe3jOsK5vOr5n4KmigJ9lbTo=
|
github.com/gaissmai/bart v0.26.0 h1:xOZ57E9hJLBiQaSyeZa9wgWhGuzfGACgqp4BE77OkO0=
|
||||||
github.com/gaissmai/bart v0.26.1/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
github.com/gaissmai/bart v0.26.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
||||||
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||||
@@ -60,8 +60,6 @@ github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX
|
|||||||
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
||||||
github.com/google/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8=
|
github.com/google/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8=
|
||||||
github.com/google/gopacket v1.1.19/go.mod h1:iJ8V8n6KS+z2U1A8pUwu8bW5SyEMkXJB8Yo/Vo+TKTo=
|
github.com/google/gopacket v1.1.19/go.mod h1:iJ8V8n6KS+z2U1A8pUwu8bW5SyEMkXJB8Yo/Vo+TKTo=
|
||||||
github.com/guptarohit/asciigraph v0.9.0 h1:MvCSRRVkT2XvU1IO6n92o7l7zqx1DiFaoszOUZQztbY=
|
|
||||||
github.com/guptarohit/asciigraph v0.9.0/go.mod h1:dYl5wwK4gNsnFf9Zp+l06rFiDZ5YtXM6x7SRWZ3KGag=
|
|
||||||
github.com/jpillora/backoff v1.0.0/go.mod h1:J/6gKK9jxlEcS3zixgDgUAsiuZ7yrSoa/FX5e0EB2j4=
|
github.com/jpillora/backoff v1.0.0/go.mod h1:J/6gKK9jxlEcS3zixgDgUAsiuZ7yrSoa/FX5e0EB2j4=
|
||||||
github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU=
|
github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU=
|
||||||
github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
|
github.com/json-iterator/go v1.1.10/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4=
|
||||||
@@ -184,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
|||||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
||||||
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
|
||||||
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg=
|
||||||
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
|
||||||
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
|
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
||||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
|
|||||||
@@ -1,57 +0,0 @@
|
|||||||
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
|
|
||||||
@@ -1,21 +0,0 @@
|
|||||||
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")
|
|
||||||
)
|
|
||||||
@@ -1,29 +0,0 @@
|
|||||||
// 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;
|
|
||||||
}
|
|
||||||
@@ -1,116 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,446 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,662 +0,0 @@
|
|||||||
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")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -1,54 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
@@ -1,63 +0,0 @@
|
|||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@@ -1,173 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,361 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
+813
@@ -0,0 +1,813 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
"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.Error("Failed to generate index",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Unable to handshake with host because no certificate is available",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", v,
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
crtHs := cs.getHandshakeBytes(v)
|
||||||
|
if crtHs == nil {
|
||||||
|
f.l.Error("Unable to handshake with host because no certificate handshake bytes is available",
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", v,
|
||||||
|
)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
ci, err := NewConnectionState(cs, crt, true, noise.HandshakeIX)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to create connection state",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", v,
|
||||||
|
)
|
||||||
|
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.Error("Failed to marshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"certVersion", v,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Failed to call noise.WriteMessage",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hh.hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Unable to handshake with host because no certificate is available",
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 0, "style": "ix_psk0"},
|
||||||
|
"certVersion", cs.initiatingVersion,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ci, err := NewConnectionState(cs, crt, false, noise.HandshakeIX)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to create connection state",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Failed to call noise.ReadMessage",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
hs := &NebulaHandshake{}
|
||||||
|
err = hs.Unmarshal(msg)
|
||||||
|
if err != nil || hs.Details == nil {
|
||||||
|
f.l.Error("Failed unmarshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
rc, err := cert.Recombine(cert.Version(hs.Details.CertVersion), hs.Details.Cert, ci.H.PeerStatic(), ci.Curve())
|
||||||
|
if err != nil {
|
||||||
|
f.l.Info("Handshake did not contain a certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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>"
|
||||||
|
}
|
||||||
|
|
||||||
|
attrs := []slog.Attr{
|
||||||
|
slog.Any("error", err),
|
||||||
|
slog.Any("from", via),
|
||||||
|
slog.Any("handshake", m{"stage": 1, "style": "ix_psk0"}),
|
||||||
|
slog.Any("certVpnNetworks", rc.Networks()),
|
||||||
|
slog.String("certFingerprint", fp),
|
||||||
|
}
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
attrs = append(attrs, slog.Any("cert", rc))
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogAttrs is intentional: attrs is a pre-built []slog.Attr slice that
|
||||||
|
// callers grow conditionally, which has no pair-form equivalent.
|
||||||
|
//nolint:sloglint
|
||||||
|
f.l.LogAttrs(context.Background(), slog.LevelInfo, "Invalid certificate from host", attrs...)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||||
|
f.l.Info("public key mismatch between certificate and handshake",
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
"cert", remoteCert,
|
||||||
|
)
|
||||||
|
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.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Might be unable to handshake with host due to missing certificate version",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
"cert", remoteCert,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// Record the certificate we are actually using
|
||||||
|
ci.myCert = myCertOtherVersion
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||||
|
f.l.Info("No networks in certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"cert", remoteCert,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Refusing to handshake with myself",
|
||||||
|
"vpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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()) {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
myIndex, err := generateIndex(f.l)
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to generate index",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.With(
|
||||||
|
"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.Error("Unable to handshake with host because no certificate handshake bytes is available",
|
||||||
|
"myCertVersion", ci.myCert.Version(),
|
||||||
|
)
|
||||||
|
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.Error("Failed to marshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Failed to call noise.WriteMessage",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
} else if dKey == nil || eKey == nil {
|
||||||
|
f.l.Error("Noise did not arrive at a key",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Error("Failed to send handshake message",
|
||||||
|
"vpnAddrs", existing.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cached", true,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
f.l.Info("Handshake message sent",
|
||||||
|
"vpnAddrs", existing.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cached", true,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
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.Info("Handshake message sent",
|
||||||
|
"vpnAddrs", existing.vpnAddrs,
|
||||||
|
"relay", via.relayHI.vpnAddrs[0],
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cached", true,
|
||||||
|
)
|
||||||
|
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.Info("Handshake too old",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"oldHandshakeTime", existing.lastHandshakeTime,
|
||||||
|
"newHandshakeTime", hostinfo.lastHandshakeTime,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.Error("Failed to add HostInfo due to localIndex collision",
|
||||||
|
"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"},
|
||||||
|
"localIndex", hostinfo.localIndexId,
|
||||||
|
"collision", existing.vpnAddrs,
|
||||||
|
)
|
||||||
|
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.Error("Failed to add HostInfo to HostMap",
|
||||||
|
"error", err,
|
||||||
|
"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"},
|
||||||
|
)
|
||||||
|
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.With(
|
||||||
|
"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": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Failed to send handshake", "error", err)
|
||||||
|
} 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.Info("Handshake message sent",
|
||||||
|
"vpnAddrs", vpnAddrs,
|
||||||
|
"relay", via.relayHI.vpnAddrs[0],
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"fingerprint", fingerprint,
|
||||||
|
"issuer", issuer,
|
||||||
|
"initiatorIndex", hs.Details.InitiatorIndex,
|
||||||
|
"responderIndex", hs.Details.ResponderIndex,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
f.connectionManager.AddTrafficWatch(hostinfo)
|
||||||
|
|
||||||
|
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
||||||
|
|
||||||
|
// Don't wait for UpdateWorker
|
||||||
|
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||||
|
f.lightHouse.TriggerUpdate()
|
||||||
|
}
|
||||||
|
|
||||||
|
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()) {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
msg, eKey, dKey, err := ci.H.ReadMessage(nil, packet[header.Len:])
|
||||||
|
if err != nil {
|
||||||
|
f.l.Error("Failed to call noise.ReadMessage",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"header", h,
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.Error("Noise did not arrive at a key",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.Error("Failed unmarshal handshake message",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.Info("Handshake did not contain a certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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>"
|
||||||
|
}
|
||||||
|
|
||||||
|
attrs := []slog.Attr{
|
||||||
|
slog.Any("error", err),
|
||||||
|
slog.Any("from", via),
|
||||||
|
slog.Any("vpnAddrs", hostinfo.vpnAddrs),
|
||||||
|
slog.Any("handshake", m{"stage": 2, "style": "ix_psk0"}),
|
||||||
|
slog.String("certFingerprint", fp),
|
||||||
|
slog.Any("certVpnNetworks", rc.Networks()),
|
||||||
|
}
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
attrs = append(attrs, slog.Any("cert", rc))
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogAttrs is intentional: attrs is a pre-built []slog.Attr slice that
|
||||||
|
// callers grow conditionally, which has no pair-form equivalent.
|
||||||
|
//nolint:sloglint
|
||||||
|
f.l.LogAttrs(context.Background(), slog.LevelInfo, "Invalid certificate from host", attrs...)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if !bytes.Equal(remoteCert.Certificate.PublicKey(), ci.H.PeerStatic()) {
|
||||||
|
f.l.Info("public key mismatch between certificate and handshake",
|
||||||
|
"from", via,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
"cert", remoteCert,
|
||||||
|
)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(remoteCert.Certificate.Networks()) == 0 {
|
||||||
|
f.l.Info("No networks in certificate",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"cert", remoteCert,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
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.Info("Incorrect host responded to handshake",
|
||||||
|
"intendedVpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"haveVpnNetworks", vpnNetworks,
|
||||||
|
"from", via,
|
||||||
|
"certName", certName,
|
||||||
|
"certVersion", certVersion,
|
||||||
|
"handshake", m{"stage": 2, "style": "ix_psk0"},
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.Info("Blocked addresses for handshakes",
|
||||||
|
"blockedUdpAddrs", newHH.hostinfo.remotes.CopyBlockedRemotes(),
|
||||||
|
"vpnNetworks", vpnNetworks,
|
||||||
|
"remotes", newHH.hostinfo.remotes.CopyAddrs(f.hostMap.GetPreferredRanges()),
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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.With(
|
||||||
|
"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": 2, "style": "ix_psk0"},
|
||||||
|
"durationNs", duration,
|
||||||
|
"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.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Sending stored packets",
|
||||||
|
"count", 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)
|
||||||
|
|
||||||
|
// Don't wait for UpdateWorker
|
||||||
|
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
||||||
|
f.lightHouse.TriggerUpdate()
|
||||||
|
}
|
||||||
|
|
||||||
|
return false
|
||||||
|
}
|
||||||
+169
-605
@@ -14,7 +14,6 @@ import (
|
|||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
@@ -23,18 +22,7 @@ const (
|
|||||||
DefaultHandshakeTryInterval = time.Millisecond * 100
|
DefaultHandshakeTryInterval = time.Millisecond * 100
|
||||||
DefaultHandshakeRetries = 10
|
DefaultHandshakeRetries = 10
|
||||||
DefaultHandshakeTriggerBuffer = 64
|
DefaultHandshakeTriggerBuffer = 64
|
||||||
|
DefaultUseRelays = true
|
||||||
// maxCachedPackets is how many unsent packets we'll buffer per pending
|
|
||||||
// handshake before dropping further ones.
|
|
||||||
maxCachedPackets = 100
|
|
||||||
|
|
||||||
// HandshakePacket map keys mirror the IX protocol stage convention:
|
|
||||||
// stage 0 = the initiator's first message (and what the responder
|
|
||||||
// receives, stripped of header)
|
|
||||||
// stage 2 = the responder's reply
|
|
||||||
// Other handshake patterns will need new keys when added.
|
|
||||||
handshakePacketStage0 uint8 = 0
|
|
||||||
handshakePacketStage2 uint8 = 2
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -42,6 +30,7 @@ var (
|
|||||||
tryInterval: DefaultHandshakeTryInterval,
|
tryInterval: DefaultHandshakeTryInterval,
|
||||||
retries: DefaultHandshakeRetries,
|
retries: DefaultHandshakeRetries,
|
||||||
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
triggerBuffer: DefaultHandshakeTriggerBuffer,
|
||||||
|
useRelays: DefaultUseRelays,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -49,6 +38,7 @@ type HandshakeConfig struct {
|
|||||||
tryInterval time.Duration
|
tryInterval time.Duration
|
||||||
retries int64
|
retries int64
|
||||||
triggerBuffer int
|
triggerBuffer int
|
||||||
|
useRelays bool
|
||||||
|
|
||||||
messageMetrics *MessageMetrics
|
messageMetrics *MessageMetrics
|
||||||
}
|
}
|
||||||
@@ -86,11 +76,10 @@ type HandshakeHostInfo struct {
|
|||||||
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
packetStore []*cachedPacket // A set of packets to be transmitted once the handshake completes
|
||||||
|
|
||||||
hostinfo *HostInfo
|
hostinfo *HostInfo
|
||||||
machine *handshake.Machine // The handshake state machine, set during stage 0 (initiator) or beginHandshake (responder multi-message)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hh *HandshakeHostInfo) cachePacket(l *slog.Logger, t header.MessageType, st header.MessageSubType, packet []byte, f packetCallback, m *cachedPacketMetrics) {
|
func (hh *HandshakeHostInfo) cachePacket(l *slog.Logger, t header.MessageType, st header.MessageSubType, packet []byte, f packetCallback, m *cachedPacketMetrics) {
|
||||||
if len(hh.packetStore) < maxCachedPackets {
|
if len(hh.packetStore) < 100 {
|
||||||
tempPacket := make([]byte, len(packet))
|
tempPacket := make([]byte, len(packet))
|
||||||
copy(tempPacket, packet)
|
copy(tempPacket, packet)
|
||||||
|
|
||||||
@@ -148,18 +137,6 @@ func (hm *HandshakeManager) Run(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *header.H) {
|
func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *header.H) {
|
||||||
// Gate on known handshake subtypes. Unknown subtypes (or future ones we
|
|
||||||
// don't yet support) are dropped here rather than silently routed through
|
|
||||||
// the IX path. Add a case when introducing a new pattern.
|
|
||||||
switch h.Subtype {
|
|
||||||
case header.HandshakeIXPSK0:
|
|
||||||
// supported
|
|
||||||
default:
|
|
||||||
hm.l.Debug("dropping handshake with unsupported subtype",
|
|
||||||
"from", via, "subtype", h.Subtype)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// First remote allow list check before we know the vpnIp
|
// First remote allow list check before we know the vpnIp
|
||||||
if !via.IsRelayed {
|
if !via.IsRelayed {
|
||||||
if !hm.lightHouse.GetRemoteAllowList().AllowUnknownVpnAddr(via.UdpAddr.Addr()) {
|
if !hm.lightHouse.GetRemoteAllowList().AllowUnknownVpnAddr(via.UdpAddr.Addr()) {
|
||||||
@@ -168,27 +145,19 @@ func (hm *HandshakeManager) HandleIncoming(via ViaSender, packet []byte, h *head
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// First message of a new handshake. The wire format requires RemoteIndex
|
switch h.Subtype {
|
||||||
// to be zero here (the initiator has no responder index to fill in yet),
|
case header.HandshakeIXPSK0:
|
||||||
// and generateIndex never allocates 0, so any non-zero RemoteIndex on a
|
switch h.MessageCounter {
|
||||||
// stage-1 packet is malformed or someone probing for an index collision.
|
case 1:
|
||||||
// Drop without paying the cost of running noise on a pending Machine.
|
ixHandshakeStage1(hm.f, via, packet, h)
|
||||||
if h.MessageCounter == 1 {
|
|
||||||
if h.RemoteIndex != 0 {
|
|
||||||
hm.l.Debug("dropping stage-1 handshake with non-zero RemoteIndex",
|
|
||||||
"from", via, "remoteIndex", h.RemoteIndex)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
hm.beginHandshake(via, packet, h)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Continuation message must match a pending handshake by index.
|
case 2:
|
||||||
// Anything else is an orphaned packet (e.g., late retransmit after
|
newHostinfo := hm.queryIndex(h.RemoteIndex)
|
||||||
// timeout) and is dropped.
|
tearDown := ixHandshakeStage2(hm.f, via, newHostinfo, packet, h)
|
||||||
if hh := hm.queryIndex(h.RemoteIndex); hh != nil {
|
if tearDown && newHostinfo != nil {
|
||||||
hm.continueHandshake(via, hh, packet)
|
hm.DeleteHostInfo(newHostinfo.hostinfo)
|
||||||
return
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -214,22 +183,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hostinfo := hh.hostinfo
|
hostinfo := hh.hostinfo
|
||||||
// If we are out of time, clean up
|
// If we are out of time, clean up
|
||||||
if hh.counter >= hm.config.retries {
|
if hh.counter >= hm.config.retries {
|
||||||
fields := []any{
|
hh.hostinfo.logger(hm.l).Info("Handshake timed out",
|
||||||
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
"udpAddrs", hh.hostinfo.remotes.CopyAddrs(hm.mainHostMap.GetPreferredRanges()),
|
||||||
"initiatorIndex", hh.hostinfo.localIndexId,
|
"initiatorIndex", hh.hostinfo.localIndexId,
|
||||||
"remoteIndex", hh.hostinfo.remoteIndexId,
|
"remoteIndex", hh.hostinfo.remoteIndexId,
|
||||||
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
"durationNs", time.Since(hh.startTime).Nanoseconds(),
|
||||||
}
|
)
|
||||||
// hh.machine can be nil here if buildStage0Packet never succeeded
|
|
||||||
// (e.g., no certificate available). In that case there's no useful
|
|
||||||
// handshake metadata to log.
|
|
||||||
if hh.machine != nil {
|
|
||||||
fields = append(fields, "handshake", m{
|
|
||||||
"stage": uint64(hh.machine.MessageIndex()),
|
|
||||||
"style": header.SubTypeName(header.Handshake, hh.machine.Subtype()),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
hh.hostinfo.logger(hm.l).Info("Handshake timed out", fields...)
|
|
||||||
hm.metricTimedOut.Inc(1)
|
hm.metricTimedOut.Inc(1)
|
||||||
hm.DeleteHostInfo(hostinfo)
|
hm.DeleteHostInfo(hostinfo)
|
||||||
return
|
return
|
||||||
@@ -240,25 +200,12 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
|
|
||||||
// Check if we have a handshake packet to transmit yet
|
// Check if we have a handshake packet to transmit yet
|
||||||
if !hh.ready {
|
if !hh.ready {
|
||||||
if !hm.buildStage0Packet(hh) {
|
if !ixHandshakeStage0(hm.f, hh) {
|
||||||
hm.OutboundHandshakeTimer.Add(vpnIp, hm.config.tryInterval*time.Duration(hh.counter))
|
hm.OutboundHandshakeTimer.Add(vpnIp, hm.config.tryInterval*time.Duration(hh.counter))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TODO: this hardcodes "always retransmit stage 0", which is correct for
|
|
||||||
// IX (the initiator only ever sends one packet, msg1) but wrong the
|
|
||||||
// moment a 3+ message pattern lands. The retry loop should resend the
|
|
||||||
// most recent outgoing message, not always stage 0. That implies
|
|
||||||
// HandshakeHostInfo tracking a single "currentOutbound" packet (bytes +
|
|
||||||
// header metadata) that gets replaced as the handshake progresses,
|
|
||||||
// instead of indexing into HandshakePacket.
|
|
||||||
stage0 := hostinfo.HandshakePacket[handshakePacketStage0]
|
|
||||||
hsFields := m{
|
|
||||||
"stage": uint64(hh.machine.MessageIndex()),
|
|
||||||
"style": header.SubTypeName(header.Handshake, hh.machine.Subtype()),
|
|
||||||
}
|
|
||||||
|
|
||||||
// Get a remotes object if we don't already have one.
|
// Get a remotes object if we don't already have one.
|
||||||
// This is mainly to protect us as this should never be the case
|
// This is mainly to protect us as this should never be the case
|
||||||
// NB ^ This comment doesn't jive. It's how the thing gets initialized.
|
// NB ^ This comment doesn't jive. It's how the thing gets initialized.
|
||||||
@@ -292,13 +239,13 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
// Send the handshake to all known ips, stage 2 takes care of assigning the hostinfo.remote based on the first to reply
|
||||||
var sentTo []netip.AddrPort
|
var sentTo []netip.AddrPort
|
||||||
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
hostinfo.remotes.ForEach(hm.mainHostMap.GetPreferredRanges(), func(addr netip.AddrPort, _ bool) {
|
||||||
hm.messageMetrics.Tx(header.Handshake, hh.machine.Subtype(), 1)
|
hm.messageMetrics.Tx(header.Handshake, header.MessageSubType(hostinfo.HandshakePacket[0][1]), 1)
|
||||||
err := hm.outside.WriteTo(stage0, addr)
|
err := hm.outside.WriteTo(hostinfo.HandshakePacket[0], addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
hostinfo.logger(hm.l).Error("Failed to send handshake message",
|
||||||
"udpAddr", addr,
|
"udpAddr", addr,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
"error", err,
|
"error", err,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -313,17 +260,156 @@ func (hm *HandshakeManager) handleOutbound(vpnIp netip.Addr, lighthouseTriggered
|
|||||||
hostinfo.logger(hm.l).Info("Handshake message sent",
|
hostinfo.logger(hm.l).Info("Handshake message sent",
|
||||||
"udpAddrs", sentTo,
|
"udpAddrs", sentTo,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
)
|
)
|
||||||
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
} else if hm.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
hostinfo.logger(hm.l).Debug("Handshake message sent",
|
||||||
"udpAddrs", sentTo,
|
"udpAddrs", sentTo,
|
||||||
"initiatorIndex", hostinfo.localIndexId,
|
"initiatorIndex", hostinfo.localIndexId,
|
||||||
"handshake", hsFields,
|
"handshake", m{"stage": 1, "style": "ix_psk0"},
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
hm.f.relayManager.StartRelays(hm.f, vpnIp, hostinfo, stage0)
|
if hm.config.useRelays && len(hostinfo.remotes.relays) > 0 {
|
||||||
|
hostinfo.logger(hm.l).Info("Attempt to relay through hosts", "relays", hostinfo.remotes.relays)
|
||||||
|
// Send a RelayRequest to all known Relay IP's
|
||||||
|
for _, relay := range hostinfo.remotes.relays {
|
||||||
|
// Don't relay through the host I'm trying to connect to
|
||||||
|
if relay == vpnIp {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Don't relay to myself
|
||||||
|
if hm.f.myVpnAddrsTable.Contains(relay) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
relayHostInfo := hm.mainHostMap.QueryVpnAddr(relay)
|
||||||
|
if relayHostInfo == nil || !relayHostInfo.remote.IsValid() {
|
||||||
|
hostinfo.logger(hm.l).Info("Establish tunnel to relay target", "relay", relay.String())
|
||||||
|
hm.f.Handshake(relay)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// Check the relay HostInfo to see if we already established a relay through
|
||||||
|
existingRelay, ok := relayHostInfo.relayState.QueryRelayForByIp(vpnIp)
|
||||||
|
if !ok {
|
||||||
|
// No relays exist or requested yet.
|
||||||
|
if relayHostInfo.remote.IsValid() {
|
||||||
|
idx, err := AddRelay(hm.l, relayHostInfo, hm.mainHostMap, vpnIp, nil, TerminalType, Requested)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(hm.l).Info("Failed to add relay to hostmap", "relay", relay.String(), "error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
m := NebulaControl{
|
||||||
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
|
InitiatorRelayIndex: idx,
|
||||||
|
}
|
||||||
|
|
||||||
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
|
case cert.Version1:
|
||||||
|
if !hm.f.myVpnAddrs[0].Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnIp.Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b := hm.f.myVpnAddrs[0].As4()
|
||||||
|
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
b = vpnIp.As4()
|
||||||
|
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
case cert.Version2:
|
||||||
|
m.RelayFromAddr = netAddrToProtoAddr(hm.f.myVpnAddrs[0])
|
||||||
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(hm.l).Error("Unknown certificate version found while creating relay")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
msg, err := m.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(hm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
|
} else {
|
||||||
|
hm.f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
|
hm.l.Info("send CreateRelayRequest",
|
||||||
|
"relayFrom", hm.f.myVpnAddrs[0],
|
||||||
|
"relayTo", vpnIp,
|
||||||
|
"initiatorRelayIndex", idx,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
switch existingRelay.State {
|
||||||
|
case Established:
|
||||||
|
hostinfo.logger(hm.l).Info("Send handshake via relay", "relay", relay.String())
|
||||||
|
hm.f.SendVia(relayHostInfo, existingRelay, hostinfo.HandshakePacket[0], make([]byte, 12), make([]byte, mtu), false)
|
||||||
|
case Disestablished:
|
||||||
|
// Mark this relay as 'requested'
|
||||||
|
relayHostInfo.relayState.UpdateRelayForByIpState(vpnIp, Requested)
|
||||||
|
fallthrough
|
||||||
|
case Requested:
|
||||||
|
hostinfo.logger(hm.l).Info("Re-send CreateRelay request", "relay", relay.String())
|
||||||
|
// Re-send the CreateRelay request, in case the previous one was lost.
|
||||||
|
m := NebulaControl{
|
||||||
|
Type: NebulaControl_CreateRelayRequest,
|
||||||
|
InitiatorRelayIndex: existingRelay.LocalIndex,
|
||||||
|
}
|
||||||
|
|
||||||
|
switch relayHostInfo.GetCert().Certificate.Version() {
|
||||||
|
case cert.Version1:
|
||||||
|
if !hm.f.myVpnAddrs[0].Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if !vpnIp.Is4() {
|
||||||
|
hostinfo.logger(hm.l).Error("can not establish v1 relay with a v6 remote network because the relay is not running a current nebula version")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
b := hm.f.myVpnAddrs[0].As4()
|
||||||
|
m.OldRelayFromAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
b = vpnIp.As4()
|
||||||
|
m.OldRelayToAddr = binary.BigEndian.Uint32(b[:])
|
||||||
|
case cert.Version2:
|
||||||
|
m.RelayFromAddr = netAddrToProtoAddr(hm.f.myVpnAddrs[0])
|
||||||
|
m.RelayToAddr = netAddrToProtoAddr(vpnIp)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(hm.l).Error("Unknown certificate version found while creating relay")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
msg, err := m.Marshal()
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(hm.l).Error("Failed to marshal Control message to create relay", "error", err)
|
||||||
|
} else {
|
||||||
|
// This must send over the hostinfo, not over hm.Hosts[ip]
|
||||||
|
hm.f.SendMessageToHostInfo(header.Control, 0, relayHostInfo, msg, make([]byte, 12), make([]byte, mtu))
|
||||||
|
hm.l.Info("send CreateRelayRequest",
|
||||||
|
"relayFrom", hm.f.myVpnAddrs[0],
|
||||||
|
"relayTo", vpnIp,
|
||||||
|
"initiatorRelayIndex", existingRelay.LocalIndex,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
case PeerRequested:
|
||||||
|
// PeerRequested only occurs in Forwarding relays, not Terminal relays, and this is a Terminal relay case.
|
||||||
|
fallthrough
|
||||||
|
default:
|
||||||
|
hostinfo.logger(hm.l).Error("Relay unexpected state",
|
||||||
|
"vpnIp", vpnIp,
|
||||||
|
"state", existingRelay.State,
|
||||||
|
"relay", relay,
|
||||||
|
)
|
||||||
|
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
// If a lighthouse triggered this attempt then we are still in the timer wheel and do not need to re-add
|
||||||
if !lighthouseTriggered {
|
if !lighthouseTriggered {
|
||||||
@@ -501,7 +587,7 @@ func (hm *HandshakeManager) Complete(hostinfo *HostInfo, f *Interface) {
|
|||||||
// allocateIndex generates a unique localIndexId for this HostInfo
|
// allocateIndex generates a unique localIndexId for this HostInfo
|
||||||
// and adds it to the pendingHostMap. Will error if we are unable to generate
|
// and adds it to the pendingHostMap. Will error if we are unable to generate
|
||||||
// a unique localIndexId
|
// a unique localIndexId
|
||||||
func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error) {
|
func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) error {
|
||||||
hm.mainHostMap.RLock()
|
hm.mainHostMap.RLock()
|
||||||
defer hm.mainHostMap.RUnlock()
|
defer hm.mainHostMap.RUnlock()
|
||||||
hm.Lock()
|
hm.Lock()
|
||||||
@@ -510,7 +596,7 @@ func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error)
|
|||||||
for range 32 {
|
for range 32 {
|
||||||
index, err := generateIndex(hm.l)
|
index, err := generateIndex(hm.l)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
_, inPending := hm.indexes[index]
|
_, inPending := hm.indexes[index]
|
||||||
@@ -519,11 +605,11 @@ func (hm *HandshakeManager) allocateIndex(hh *HandshakeHostInfo) (uint32, error)
|
|||||||
if !inMain && !inPending {
|
if !inMain && !inPending {
|
||||||
hh.hostinfo.localIndexId = index
|
hh.hostinfo.localIndexId = index
|
||||||
hm.indexes[index] = hh
|
hm.indexes[index] = hh
|
||||||
return index, nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return 0, errors.New("failed to generate unique localIndexId")
|
return errors.New("failed to generate unique localIndexId")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
func (hm *HandshakeManager) DeleteHostInfo(hostinfo *HostInfo) {
|
||||||
@@ -642,525 +728,3 @@ func generateIndex(l *slog.Logger) (uint32, error) {
|
|||||||
func hsTimeout(tries int64, interval time.Duration) time.Duration {
|
func hsTimeout(tries int64, interval time.Duration) time.Duration {
|
||||||
return time.Duration(tries / 2 * ((2 * int64(interval)) + (tries-1)*int64(interval)))
|
return time.Duration(tries / 2 * ((2 * int64(interval)) + (tries-1)*int64(interval)))
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildStage0Packet creates the initial handshake packet for the initiator.
|
|
||||||
func (hm *HandshakeManager) buildStage0Packet(hh *HandshakeHostInfo) bool {
|
|
||||||
cs := hm.f.pki.getCertState()
|
|
||||||
v := cs.DefaultVersion()
|
|
||||||
if hh.initiatingVersionOverride != cert.VersionPre1 {
|
|
||||||
v = hh.initiatingVersionOverride
|
|
||||||
} else if v < cert.Version2 {
|
|
||||||
for _, a := range hh.hostinfo.vpnAddrs {
|
|
||||||
if a.Is6() {
|
|
||||||
v = cert.Version2
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
cred := cs.GetCredential(v)
|
|
||||||
if cred == nil {
|
|
||||||
hm.f.l.Error("Unable to handshake with host because no certificate is available",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs, "certVersion", v)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
machine, err := handshake.NewMachine(
|
|
||||||
v, cs.GetCredential,
|
|
||||||
hm.certVerifier(), func() (uint32, error) { return hm.allocateIndex(hh) },
|
|
||||||
true, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
hm.f.l.Error("Failed to create handshake machine",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs, "error", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
msg, err := machine.Initiate(nil)
|
|
||||||
if err != nil {
|
|
||||||
hm.f.l.Error("Failed to initiate handshake",
|
|
||||||
"vpnAddrs", hh.hostinfo.vpnAddrs, "error", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
// hostinfo.ConnectionState stays nil until the handshake completes in
|
|
||||||
// continueHandshake. Pre-completion control surfaces guard with nil
|
|
||||||
// checks; the data plane never observes a pending hostinfo.
|
|
||||||
hh.hostinfo.HandshakePacket[handshakePacketStage0] = msg
|
|
||||||
hh.machine = machine
|
|
||||||
hh.ready = true
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// beginHandshake handles an incoming handshake packet that doesn't match any
|
|
||||||
// existing pending handshake. It creates a new responder Machine and processes
|
|
||||||
// the first message.
|
|
||||||
func (hm *HandshakeManager) beginHandshake(via ViaSender, packet []byte, h *header.H) {
|
|
||||||
f := hm.f
|
|
||||||
cs := f.pki.getCertState()
|
|
||||||
|
|
||||||
v := cs.DefaultVersion()
|
|
||||||
if cs.GetCredential(v) == nil {
|
|
||||||
f.l.Error("Unable to handshake with host because no certificate is available",
|
|
||||||
"from", via, "certVersion", v)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
machine, err := handshake.NewMachine(
|
|
||||||
v, cs.GetCredential,
|
|
||||||
hm.certVerifier(), func() (uint32, error) { return generateIndex(f.l) },
|
|
||||||
false, header.HandshakeIXPSK0,
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to create handshake machine", "from", via, "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
response, result, err := machine.ProcessPacket(nil, packet)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to process handshake packet", "from", via, "error", err)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if result == nil {
|
|
||||||
// Multi-message pattern: the responder Machine would need to be
|
|
||||||
// registered in hm.indexes so a future inbound packet finds it via
|
|
||||||
// continueHandshake. The current manager doesn't do that yet, so
|
|
||||||
// fail loudly rather than silently dropping the in-flight handshake.
|
|
||||||
// TODO: support multi-message responder flows (XX, pqIX, etc.).
|
|
||||||
// See also the IX-shaped cipher key assignment in handshake.Machine.
|
|
||||||
f.l.Error("multi-message handshake responder is not supported",
|
|
||||||
"from", via, "error", handshake.ErrMultiMessageUnsupported)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
remoteCert := result.RemoteCert
|
|
||||||
if remoteCert == nil {
|
|
||||||
f.l.Error("Handshake did not produce a peer certificate", "from", via)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Validate peer identity
|
|
||||||
vpnAddrs, anyVpnAddrsInCommon, ok := hm.validatePeerCert(via, remoteCert)
|
|
||||||
if !ok {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := &HostInfo{
|
|
||||||
ConnectionState: newConnectionStateFromResult(result),
|
|
||||||
localIndexId: result.LocalIndex,
|
|
||||||
remoteIndexId: result.RemoteIndex,
|
|
||||||
vpnAddrs: vpnAddrs,
|
|
||||||
HandshakePacket: make(map[uint8][]byte, 0),
|
|
||||||
lastHandshakeTime: result.HandshakeTime,
|
|
||||||
relayState: RelayState{
|
|
||||||
relays: nil,
|
|
||||||
relayForByAddr: map[netip.Addr]*Relay{},
|
|
||||||
relayForByIdx: map[uint32]*Relay{},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
msg := "Handshake message received"
|
|
||||||
if !anyVpnAddrsInCommon {
|
|
||||||
msg = "Handshake message received, but no vpnNetworks in common."
|
|
||||||
}
|
|
||||||
f.l.Info(msg,
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"initiatorIndex", result.RemoteIndex,
|
|
||||||
"responderIndex", result.LocalIndex,
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
)
|
|
||||||
|
|
||||||
// packet aliases the listener's incoming buffer, so this copy must stay.
|
|
||||||
hostinfo.HandshakePacket[handshakePacketStage0] = make([]byte, len(packet[header.Len:]))
|
|
||||||
copy(hostinfo.HandshakePacket[handshakePacketStage0], packet[header.Len:])
|
|
||||||
|
|
||||||
// response was freshly allocated by ProcessPacket; safe to retain directly.
|
|
||||||
if response != nil {
|
|
||||||
hostinfo.HandshakePacket[handshakePacketStage2] = response
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo.remotes = f.lightHouse.QueryCache(vpnAddrs)
|
|
||||||
if !via.IsRelayed {
|
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
|
||||||
}
|
|
||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
|
||||||
|
|
||||||
existing, err := hm.CheckAndComplete(hostinfo, handshakePacketStage0, f)
|
|
||||||
if err != nil {
|
|
||||||
hm.handleCheckAndCompleteError(err, existing, hostinfo, via)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
|
||||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
|
||||||
hostinfo.remotes.RefreshFromHandshake(vpnAddrs)
|
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
|
||||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
|
||||||
f.lightHouse.TriggerUpdate()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// continueHandshake feeds an incoming packet to an existing pending handshake Machine.
|
|
||||||
func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostInfo, packet []byte) {
|
|
||||||
f := hm.f
|
|
||||||
|
|
||||||
hh.Lock()
|
|
||||||
defer hh.Unlock()
|
|
||||||
|
|
||||||
// Re-verify hh is still tracked. Between queryIndex returning and us taking
|
|
||||||
// hh.Lock, handleOutbound may have timed out and deleted it. Once we hold
|
|
||||||
// hh.Lock no other deleter can race our index: handleOutbound also takes
|
|
||||||
// hh.Lock first, and handleRecvError targets a main-hostmap entry with a
|
|
||||||
// different localIndexId.
|
|
||||||
hm.RLock()
|
|
||||||
cur, ok := hm.indexes[hh.hostinfo.localIndexId]
|
|
||||||
hm.RUnlock()
|
|
||||||
if !ok || cur != hh {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
hostinfo := hh.hostinfo
|
|
||||||
if !via.IsRelayed {
|
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(hostinfo.vpnAddrs, via.UdpAddr.Addr()) {
|
|
||||||
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
machine := hh.machine
|
|
||||||
if machine == nil {
|
|
||||||
f.l.Error("No handshake machine available for continuation",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
response, result, err := machine.ProcessPacket(nil, packet)
|
|
||||||
if err != nil {
|
|
||||||
// Recoverable errors are routine noise, log at Debug. Fatal errors get a Warn.
|
|
||||||
if machine.Failed() {
|
|
||||||
f.l.Warn("Failed to process handshake packet, abandoning",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via, "error", err)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
} else {
|
|
||||||
f.l.Debug("Failed to process handshake packet",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via, "error", err)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if response != nil {
|
|
||||||
hm.sendHandshakeResponse(via, response, hostinfo, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
if result == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// Handshake complete; build the ConnectionState now that we have keys and a verified peer cert.
|
|
||||||
hostinfo.ConnectionState = newConnectionStateFromResult(result)
|
|
||||||
|
|
||||||
remoteCert := result.RemoteCert
|
|
||||||
if remoteCert == nil {
|
|
||||||
f.l.Error("Handshake completed without peer certificate",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs, "from", via)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
vpnNetworks := remoteCert.Certificate.Networks()
|
|
||||||
hostinfo.remoteIndexId = result.RemoteIndex
|
|
||||||
hostinfo.lastHandshakeTime = result.HandshakeTime
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
hostinfo.SetRemote(via.UdpAddr)
|
|
||||||
} else {
|
|
||||||
hostinfo.relayState.InsertRelayTo(via.relayHI.vpnAddrs[0])
|
|
||||||
}
|
|
||||||
|
|
||||||
// Verify correct host responded (initiator check)
|
|
||||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
|
||||||
correctHostResponded := false
|
|
||||||
anyVpnAddrsInCommon := false
|
|
||||||
for i, network := range vpnNetworks {
|
|
||||||
// inside.go drops self-routed packets at the firewall stage, but we'd
|
|
||||||
// rather not let a self-handshake complete in the first place: it
|
|
||||||
// wastes a hostmap slot, suppresses no log, and obscures routing
|
|
||||||
// misconfig. Explicit refusal here mirrors the responder-side check
|
|
||||||
// in validatePeerCert.
|
|
||||||
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
|
||||||
f.l.Error("Refusing to handshake with myself",
|
|
||||||
"vpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
)
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
vpnAddrs[i] = network.Addr()
|
|
||||||
if hostinfo.vpnAddrs[0] == network.Addr() {
|
|
||||||
correctHostResponded = true
|
|
||||||
}
|
|
||||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
|
||||||
anyVpnAddrsInCommon = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !correctHostResponded {
|
|
||||||
f.l.Info("Incorrect host responded to handshake",
|
|
||||||
"intendedVpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"haveVpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
)
|
|
||||||
|
|
||||||
hm.DeleteHostInfo(hostinfo)
|
|
||||||
hm.StartHandshake(hostinfo.vpnAddrs[0], func(newHH *HandshakeHostInfo) {
|
|
||||||
newHH.hostinfo.remotes = hostinfo.remotes
|
|
||||||
newHH.hostinfo.remotes.BlockRemote(via)
|
|
||||||
newHH.packetStore = hh.packetStore
|
|
||||||
hh.packetStore = []*cachedPacket{}
|
|
||||||
hostinfo.vpnAddrs = vpnAddrs
|
|
||||||
f.sendCloseTunnel(hostinfo)
|
|
||||||
})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
duration := time.Since(hh.startTime).Nanoseconds()
|
|
||||||
msg := "Handshake message received"
|
|
||||||
if !anyVpnAddrsInCommon {
|
|
||||||
msg = "Handshake message received, but no vpnNetworks in common."
|
|
||||||
}
|
|
||||||
f.l.Info(msg,
|
|
||||||
"vpnAddrs", vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
"initiatorIndex", result.LocalIndex,
|
|
||||||
"responderIndex", result.RemoteIndex,
|
|
||||||
"handshake", m{"stage": uint64(machine.MessageIndex()), "style": header.SubTypeName(header.Handshake, machine.Subtype())},
|
|
||||||
"durationNs", duration,
|
|
||||||
"sentCachedPackets", len(hh.packetStore),
|
|
||||||
)
|
|
||||||
|
|
||||||
hostinfo.vpnAddrs = vpnAddrs
|
|
||||||
hostinfo.buildNetworks(f.myVpnNetworksTable, remoteCert.Certificate)
|
|
||||||
|
|
||||||
hm.Complete(hostinfo, f)
|
|
||||||
f.connectionManager.AddTrafficWatch(hostinfo)
|
|
||||||
|
|
||||||
if len(hh.packetStore) > 0 {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
hostinfo.logger(f.l).Debug("Sending stored packets", "count", len(hh.packetStore))
|
|
||||||
}
|
|
||||||
nb := make([]byte, 12, 12)
|
|
||||||
out := make([]byte, mtu)
|
|
||||||
for _, cp := range hh.packetStore {
|
|
||||||
//todo use a sendbatcher
|
|
||||||
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)
|
|
||||||
|
|
||||||
// Don't wait for UpdateWorker
|
|
||||||
if f.lightHouse.IsAnyLighthouseAddr(vpnAddrs) {
|
|
||||||
f.lightHouse.TriggerUpdate()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// validatePeerCert checks the peer certificate for self-connection and remote allow list.
|
|
||||||
// Returns the VPN addrs, whether any of them fall within one of our own VPN
|
|
||||||
// networks, and true if valid; false if rejected.
|
|
||||||
func (hm *HandshakeManager) validatePeerCert(via ViaSender, remoteCert *cert.CachedCertificate) ([]netip.Addr, bool, bool) {
|
|
||||||
f := hm.f
|
|
||||||
vpnNetworks := remoteCert.Certificate.Networks()
|
|
||||||
|
|
||||||
// The cert package rejects host certs with no networks at parse time, so
|
|
||||||
// reaching this state would mean an invariant was bypassed elsewhere.
|
|
||||||
// Refuse explicitly so downstream code (which indexes vpnAddrs[0]) can't
|
|
||||||
// panic if that invariant ever changes.
|
|
||||||
if len(vpnNetworks) == 0 {
|
|
||||||
f.l.Info("No networks in certificate",
|
|
||||||
"from", via, "cert", remoteCert)
|
|
||||||
return nil, false, false
|
|
||||||
}
|
|
||||||
|
|
||||||
vpnAddrs := make([]netip.Addr, len(vpnNetworks))
|
|
||||||
anyVpnAddrsInCommon := false
|
|
||||||
|
|
||||||
for i, network := range vpnNetworks {
|
|
||||||
if f.myVpnAddrsTable.Contains(network.Addr()) {
|
|
||||||
f.l.Error("Refusing to handshake with myself",
|
|
||||||
"vpnNetworks", vpnNetworks,
|
|
||||||
"from", via,
|
|
||||||
"certName", remoteCert.Certificate.Name(),
|
|
||||||
"certVersion", remoteCert.Certificate.Version(),
|
|
||||||
"fingerprint", remoteCert.Fingerprint,
|
|
||||||
"issuer", remoteCert.Certificate.Issuer(),
|
|
||||||
)
|
|
||||||
return nil, false, false
|
|
||||||
}
|
|
||||||
vpnAddrs[i] = network.Addr()
|
|
||||||
if f.myVpnNetworksTable.Contains(network.Addr()) {
|
|
||||||
anyVpnAddrsInCommon = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
if !f.lightHouse.GetRemoteAllowList().AllowAll(vpnAddrs, via.UdpAddr.Addr()) {
|
|
||||||
f.l.Debug("lighthouse.remote_allow_list denied incoming handshake",
|
|
||||||
"vpnAddrs", vpnAddrs, "from", via)
|
|
||||||
return nil, false, false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return vpnAddrs, anyVpnAddrsInCommon, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendHandshakeResponse sends a handshake response via the appropriate transport.
|
|
||||||
// cached is true when msg is a stored response being retransmitted because
|
|
||||||
// the peer's stage-1 retransmit landed (the ErrAlreadySeen path); false on a
|
|
||||||
// fresh response.
|
|
||||||
func (hm *HandshakeManager) sendHandshakeResponse(via ViaSender, msg []byte, hostinfo *HostInfo, cached bool) {
|
|
||||||
if msg == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f := hm.f
|
|
||||||
f.messageMetrics.Tx(header.Handshake, header.MessageSubType(msg[1]), 1)
|
|
||||||
|
|
||||||
// Common log fields. peerCert may be nil during intermediate
|
|
||||||
// multi-message flows (handshake hasn't completed yet); skip the cert
|
|
||||||
// block if so.
|
|
||||||
logFields := []any{
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"handshake", m{"stage": uint64(2), "style": header.SubTypeName(header.Handshake, header.HandshakeIXPSK0)},
|
|
||||||
"cached", cached,
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
}
|
|
||||||
if peerCert := hostinfo.ConnectionState.peerCert; peerCert != nil {
|
|
||||||
logFields = append(logFields,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !via.IsRelayed {
|
|
||||||
fields := append(logFields, "from", via)
|
|
||||||
err := f.outside.WriteTo(msg, via.UdpAddr)
|
|
||||||
if err != nil {
|
|
||||||
f.l.Error("Failed to send handshake message", append(fields, "error", err)...)
|
|
||||||
} else {
|
|
||||||
f.l.Info("Handshake message sent", fields...)
|
|
||||||
}
|
|
||||||
} 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])
|
|
||||||
// We received a valid handshake on this relay, so make sure the relay
|
|
||||||
// state reflects that, in case it had been marked Disestablished.
|
|
||||||
via.relayHI.relayState.UpdateRelayForByIdxState(via.remoteIdx, Established)
|
|
||||||
f.SendVia(via.relayHI, via.relay, msg, make([]byte, 12), make([]byte, mtu), false)
|
|
||||||
f.l.Info("Handshake message sent", append(logFields, "relay", via.relayHI.vpnAddrs[0])...)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// handleCheckAndCompleteError handles errors from CheckAndComplete.
|
|
||||||
// This only fires from the responder-side beginHandshake path, after the
|
|
||||||
// peer cert has been validated and ConnectionState populated, so peerCert
|
|
||||||
// is always non-nil for the cases that log it.
|
|
||||||
func (hm *HandshakeManager) handleCheckAndCompleteError(err error, existing, hostinfo *HostInfo, via ViaSender) {
|
|
||||||
f := hm.f
|
|
||||||
peerCert := hostinfo.ConnectionState.peerCert
|
|
||||||
hsFields := m{"stage": uint64(1), "style": header.SubTypeName(header.Handshake, header.HandshakeIXPSK0)}
|
|
||||||
|
|
||||||
switch err {
|
|
||||||
case ErrAlreadySeen:
|
|
||||||
if existing.SetRemoteIfPreferred(f.hostMap, via) {
|
|
||||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
|
||||||
}
|
|
||||||
// Resend the original response. The peer is committed to that response's
|
|
||||||
// ephemeral keys; a freshly-built one would have different keys and break
|
|
||||||
// the tunnel even though both sides "completed" the handshake.
|
|
||||||
if msg := existing.HandshakePacket[handshakePacketStage2]; msg != nil {
|
|
||||||
hm.sendHandshakeResponse(via, msg, existing, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
case ErrExistingHostInfo:
|
|
||||||
f.l.Info("Handshake too old",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
"oldHandshakeTime", existing.lastHandshakeTime,
|
|
||||||
"newHandshakeTime", hostinfo.lastHandshakeTime,
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
"handshake", hsFields,
|
|
||||||
)
|
|
||||||
f.SendMessageToVpnAddr(header.Test, header.TestRequest, hostinfo.vpnAddrs[0], []byte(""), make([]byte, 12, 12), make([]byte, mtu))
|
|
||||||
|
|
||||||
case ErrLocalIndexCollision:
|
|
||||||
f.l.Error("Failed to add HostInfo due to localIndex collision",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
"localIndex", hostinfo.localIndexId,
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
"handshake", hsFields,
|
|
||||||
)
|
|
||||||
|
|
||||||
default:
|
|
||||||
f.l.Error("Failed to add HostInfo to HostMap",
|
|
||||||
"vpnAddrs", hostinfo.vpnAddrs,
|
|
||||||
"from", via,
|
|
||||||
"error", err,
|
|
||||||
"certName", peerCert.Certificate.Name(),
|
|
||||||
"certVersion", peerCert.Certificate.Version(),
|
|
||||||
"fingerprint", peerCert.Fingerprint,
|
|
||||||
"issuer", peerCert.Certificate.Issuer(),
|
|
||||||
"initiatorIndex", hostinfo.remoteIndexId,
|
|
||||||
"responderIndex", hostinfo.localIndexId,
|
|
||||||
"handshake", hsFields,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// certVerifier returns a CertVerifier that validates certs against the current CA pool.
|
|
||||||
func (hm *HandshakeManager) certVerifier() handshake.CertVerifier {
|
|
||||||
return func(c cert.Certificate) (*cert.CachedCertificate, error) {
|
|
||||||
return hm.f.pki.GetCAPool().VerifyCertificate(time.Now(), c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+1
-136
@@ -5,7 +5,6 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/test"
|
"github.com/slackhq/nebula/test"
|
||||||
@@ -28,7 +27,7 @@ func Test_NewHandshakeManagerVpnIp(t *testing.T) {
|
|||||||
initiatingVersion: cert.Version1,
|
initiatingVersion: cert.Version1,
|
||||||
privateKey: []byte{},
|
privateKey: []byte{},
|
||||||
v1Cert: &dummyCert{version: cert.Version1},
|
v1Cert: &dummyCert{version: cert.Version1},
|
||||||
v1Credential: nil,
|
v1HandshakeBytes: []byte{},
|
||||||
}
|
}
|
||||||
|
|
||||||
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
blah := NewHandshakeManager(l, mainHM, lh, &udp.NoopConn{}, defaultHandshakeConfig)
|
||||||
@@ -101,137 +100,3 @@ func (mw *mockEncWriter) GetHostInfo(_ netip.Addr) *HostInfo {
|
|||||||
func (mw *mockEncWriter) GetCertState() *CertState {
|
func (mw *mockEncWriter) GetCertState() *CertState {
|
||||||
return &CertState{initiatingVersion: cert.Version2}
|
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,10 +174,6 @@ func (h *H) SubTypeName() string {
|
|||||||
return SubTypeName(h.Type, h.Subtype)
|
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
|
// SubTypeName will transform a nebula message sub type into a human string
|
||||||
func SubTypeName(t MessageType, s MessageSubType) string {
|
func SubTypeName(t MessageType, s MessageSubType) string {
|
||||||
if n, ok := subTypeMap[t]; ok {
|
if n, ok := subTypeMap[t]; ok {
|
||||||
@@ -189,16 +185,6 @@ func SubTypeName(t MessageType, s MessageSubType) string {
|
|||||||
return "unknown"
|
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
|
// NewHeader turns bytes into a header
|
||||||
func NewHeader(b []byte) (*H, error) {
|
func NewHeader(b []byte) (*H, error) {
|
||||||
h := new(H)
|
h := new(H)
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -11,25 +10,12 @@ import (
|
|||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
"github.com/slackhq/nebula/overlay/batch"
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packet, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(packet []byte, 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
|
err := newPacket(packet, false, fwPacket)
|
||||||
// only valid until the next Read on that queue. Every consumer below
|
if err != nil {
|
||||||
// (parse, self-forward, handshake cache, sendInsideMessage) reads it
|
|
||||||
// synchronously; do not retain pkt outside this call. If a future
|
|
||||||
// caller needs to keep the packet, use pkt.Clone() to detach it from
|
|
||||||
// the borrow.
|
|
||||||
//
|
|
||||||
// pkt.Bytes is either one IP datagram (GSO zero) or a TSO/USO
|
|
||||||
// superpacket. In both cases the L3+L4 headers at the start describe
|
|
||||||
// the same 5-tuple every segment will share, so a single parse +
|
|
||||||
// firewall check covers the whole superpacket.
|
|
||||||
packet := pkt.Bytes
|
|
||||||
var parsed batch.RxParsed
|
|
||||||
if err := batch.ParsePacket(packet, false, &parsed); err != nil {
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("Error while validating outbound packet",
|
f.l.Debug("Error while validating outbound packet",
|
||||||
"packet", packet,
|
"packet", packet,
|
||||||
@@ -39,8 +25,6 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
parsed.Key.Hydrate(fwPacket)
|
|
||||||
|
|
||||||
// Ignore local broadcast packets
|
// Ignore local broadcast packets
|
||||||
if f.dropLocalBroadcast {
|
if f.dropLocalBroadcast {
|
||||||
if f.myBroadcastAddrsTable.Contains(fwPacket.RemoteAddr) {
|
if f.myBroadcastAddrsTable.Contains(fwPacket.RemoteAddr) {
|
||||||
@@ -54,14 +38,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
// Write copies into the kernel queue synchronously, so seg's lifetime ends at return.
|
_, err := f.readers[q].Write(packet)
|
||||||
// A self-forwarded superpacket would be re-handed to the
|
|
||||||
// kernel as one giant blob; segment first so the loopback
|
|
||||||
// path sees one IP datagram per Write.
|
|
||||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
|
||||||
_, werr := f.readers[q].Write(seg)
|
|
||||||
return werr
|
|
||||||
})
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -77,19 +54,7 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
|
||||||
// so retaining segments past the loop is safe.
|
|
||||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
|
||||||
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
|
||||||
f.l.Debug("Failed to segment superpacket for handshake cache",
|
|
||||||
"error", err,
|
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
@@ -107,9 +72,9 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(parsed.Key, fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
dropReason := f.firewall.Drop(*fwPacket, false, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason == nil {
|
if dropReason == nil {
|
||||||
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch, rejectBuf, q)
|
f.sendInsideMessage(hostinfo, packet, nb, sendBatch, rejectBuf, q)
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, rejectBuf, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -121,7 +86,31 @@ func (f *Interface) consumeInsidePacket(pkt tio.Packet, fwPacket *firewall.Packe
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
// sendInsideMessage encrypts a firewall-approved inside packet into the
|
||||||
|
// caller's batch slot for later sendmmsg flush. When hostinfo.remote is not
|
||||||
|
// valid we fall through to the relay slow path via the unbatched sendNoMetrics
|
||||||
|
// so relay behavior is unchanged.
|
||||||
|
func (f *Interface) sendInsideMessage(hostinfo *HostInfo, p, nb []byte, sendBatch batch.TxBatcher, rejectBuf []byte, q int) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hostinfo.remote.IsValid() {
|
||||||
|
// Slow path: relay fallback. Reuse rejectBuf as the ciphertext
|
||||||
|
// scratch; sendNoMetrics arranges header space for SendVia.
|
||||||
|
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, p, nb, rejectBuf, q)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
scratch := sendBatch.Next()
|
||||||
|
if scratch == nil {
|
||||||
|
// Batch full: bypass batching and send this packet directly so we
|
||||||
|
// never drop traffic on over-subscribed iterations.
|
||||||
|
f.sendNoMetrics(header.Message, 0, ci, hostinfo, netip.AddrPort{}, p, nb, rejectBuf, q)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
ci.writeLock.Lock()
|
ci.writeLock.Lock()
|
||||||
}
|
}
|
||||||
@@ -130,38 +119,6 @@ func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, s
|
|||||||
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
f.connectionManager.Out(hostinfo)
|
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 {
|
if hostinfo.lastRebindCount != f.rebindCount {
|
||||||
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||||
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
@@ -174,94 +131,20 @@ func (f *Interface) sendInsideMessage(hostinfo *HostInfo, pkt tio.Packet, nb []b
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if !hostinfo.remote.IsValid() { //the relay path
|
out, err := ci.eKey.EncryptDanger(out, out, p, c, nb)
|
||||||
//first, find our relay hostinfo:
|
if noiseutil.EncryptLockNeeded {
|
||||||
var relayHostInfo *HostInfo
|
ci.writeLock.Unlock()
|
||||||
var relay *Relay
|
}
|
||||||
var err error
|
if err != nil {
|
||||||
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||||
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
"error", err,
|
||||||
if err != nil {
|
"udpAddr", hostinfo.remote,
|
||||||
hostinfo.relayState.DeleteRelay(relayIP)
|
"counter", c,
|
||||||
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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
err := tio.SegmentSuperpacket(pkt, func(seg []byte) error {
|
sendBatch.Commit(len(out), hostinfo.remote)
|
||||||
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
|
||||||
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
|
||||||
|
|
||||||
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
|
||||||
if out == nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var ecn byte
|
|
||||||
if ecnEnabled {
|
|
||||||
ecn = innerECN(seg)
|
|
||||||
}
|
|
||||||
sendBatch.Commit(out, hostinfo.remote, ecn)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
|
||||||
"error", err,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// innerECN returns the 2-bit IP-level ECN codepoint of an inner IPv4 or IPv6
|
|
||||||
// packet, or 0 if pkt is too short or its IP version is unrecognized. Used at
|
|
||||||
// encap to copy the inner codepoint onto the outer carrier per RFC 6040.
|
|
||||||
func innerECN(pkt []byte) byte {
|
|
||||||
if len(pkt) < 2 {
|
|
||||||
return 0
|
|
||||||
}
|
|
||||||
switch pkt[0] >> 4 {
|
|
||||||
case 4:
|
|
||||||
return pkt[1] & 0x03
|
|
||||||
case 6:
|
|
||||||
return (pkt[1] >> 4) & 0x03
|
|
||||||
}
|
|
||||||
return 0
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
@@ -394,16 +277,15 @@ func (f *Interface) getOrHandshakeConsiderRouting(fwPacket *firewall.Packet, cac
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
func (f *Interface) sendMessageNow(t header.MessageType, st header.MessageSubType, hostinfo *HostInfo, p, nb, out []byte) {
|
||||||
var parsed batch.RxParsed
|
fp := &firewall.Packet{}
|
||||||
if err := batch.ParsePacket(p, false, &parsed); err != nil {
|
err := newPacket(p, false, fp)
|
||||||
|
if err != nil {
|
||||||
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
f.l.Warn("error while parsing outgoing packet for firewall check", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
fp := &firewall.Packet{}
|
|
||||||
parsed.Key.Hydrate(fp)
|
|
||||||
|
|
||||||
// check if packet is in outbound fw rules
|
// check if packet is in outbound fw rules
|
||||||
dropReason := f.firewall.Drop(parsed.Key, fp, false, hostinfo, f.pki.GetCAPool(), nil)
|
dropReason := f.firewall.Drop(*fp, false, hostinfo, f.pki.GetCAPool(), nil)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping cached packet",
|
f.l.Debug("dropping cached packet",
|
||||||
@@ -454,13 +336,21 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
|||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) prepareSendVia(via *HostInfo,
|
// 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,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) ([]byte, error) {
|
) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
@@ -482,7 +372,7 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return nil, io.ErrShortBuffer
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -502,32 +392,13 @@ func (f *Interface) prepareSendVia(via *HostInfo,
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
return nil, err
|
return
|
||||||
}
|
}
|
||||||
f.connectionManager.RelayUsed(relay.LocalIndex)
|
err = f.writers[0].WriteTo(out, via.remote)
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
|
||||||
// handshake messages to peers through relay tunnels.
|
|
||||||
// via is the HostInfo through which the message is relayed.
|
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
|
||||||
// q indicates which writer to use to send the packet.
|
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
|
||||||
relay *Relay,
|
|
||||||
ad,
|
|
||||||
nb,
|
|
||||||
out []byte,
|
|
||||||
nocopy bool,
|
|
||||||
) {
|
|
||||||
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
|
||||||
err = f.writers[0].WriteTo(toSend, via.remote)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to WriteTo in sendVia", "error", err)
|
||||||
}
|
}
|
||||||
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
func (f *Interface) sendNoMetrics(t header.MessageType, st header.MessageSubType, ci *ConnectionState, hostinfo *HostInfo, remote netip.AddrPort, p, nb, out []byte, q int) {
|
||||||
|
|||||||
+23
-72
@@ -6,14 +6,12 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/util"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
@@ -50,14 +48,7 @@ type InterfaceConfig struct {
|
|||||||
reQueryWait time.Duration
|
reQueryWait time.Duration
|
||||||
|
|
||||||
ConntrackCacheTimeout time.Duration
|
ConntrackCacheTimeout time.Duration
|
||||||
|
l *slog.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 {
|
type Interface struct {
|
||||||
@@ -81,16 +72,7 @@ type Interface struct {
|
|||||||
routines int
|
routines int
|
||||||
disconnectInvalid atomic.Bool
|
disconnectInvalid atomic.Bool
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
// cpuAffinity, when non-empty, names the CPUs each TUN reader goroutine
|
relayManager *relayManager
|
||||||
// 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
|
tryPromoteEvery atomic.Uint32
|
||||||
reQueryEvery atomic.Uint32
|
reQueryEvery atomic.Uint32
|
||||||
@@ -220,7 +202,6 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
relayManager: c.relayManager,
|
relayManager: c.relayManager,
|
||||||
connectionManager: c.connectionManager,
|
connectionManager: c.connectionManager,
|
||||||
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
conntrackCacheTimeout: c.ConntrackCacheTimeout,
|
||||||
cpuAffinity: c.CpuAffinity,
|
|
||||||
|
|
||||||
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
metricHandshakes: metrics.GetOrRegisterHistogram("handshakes", nil, metrics.NewExpDecaySample(1028, 0.015)),
|
||||||
messageMetrics: c.MessageMetrics,
|
messageMetrics: c.MessageMetrics,
|
||||||
@@ -279,16 +260,7 @@ func (f *Interface) activate() error {
|
|||||||
}
|
}
|
||||||
f.readers = f.inside.Readers()
|
f.readers = f.inside.Readers()
|
||||||
for i := range f.readers {
|
for i := range f.readers {
|
||||||
caps := tio.QueueCapabilities(f.readers[i])
|
f.batchers[i] = batch.NewTCPCoalescer(f.readers[i])
|
||||||
if caps.TSO || caps.USO {
|
|
||||||
// Multi-lane: TCP gets coalesced when TSO is on, UDP when USO
|
|
||||||
// is on, everything else (and either lane disabled) falls
|
|
||||||
// through to passthrough so non-IP / non-TCP-UDP traffic still
|
|
||||||
// reaches the TUN.
|
|
||||||
f.batchers[i] = batch.NewMultiCoalescer(f.readers[i], caps.TSO, caps.USO)
|
|
||||||
} else {
|
|
||||||
f.batchers[i] = batch.NewPassthrough(f.readers[i])
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
f.wg.Add(1) // for us to wait on Close() to return
|
f.wg.Add(1) // for us to wait on Close() to return
|
||||||
@@ -348,16 +320,17 @@ func (f *Interface) listenOut(i int) {
|
|||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
parsedRx := &batch.RxParsed{}
|
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
coalescer := f.batchers[i]
|
||||||
|
|
||||||
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte) {
|
||||||
plaintext := f.batchers[i].Reserve(len(payload))
|
plaintext := f.batchers[i].Reserve(len(payload))
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, parsedRx, lhh, nb, i, ctCache.Get(), meta)
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
}
|
}
|
||||||
|
|
||||||
flusher := func() {
|
flusher := func() {
|
||||||
if err := f.batchers[i].Flush(); err != nil {
|
if err := coalescer.Flush(); err != nil {
|
||||||
f.l.Error("Failed to flush tun coalescer", "error", err)
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -373,27 +346,8 @@ func (f *Interface) listenOut(i int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader tio.Queue, i int) {
|
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)
|
rejectBuf := make([]byte, mtu)
|
||||||
sb := batch.NewSendBatch(f.writers[i], batch.SendBatchCap, udp.MTU+32)
|
sb := batch.NewSendBatch(batch.SendBatchCap, udp.MTU+32)
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
@@ -409,24 +363,35 @@ func (f *Interface) listenIn(reader tio.Queue, i int) {
|
|||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
|
sb.Reset()
|
||||||
for _, pkt := range pkts {
|
for _, pkt := range pkts {
|
||||||
|
if sb.Len() >= sb.Cap() {
|
||||||
|
f.flushBatch(sb, i)
|
||||||
|
sb.Reset()
|
||||||
|
}
|
||||||
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
f.consumeInsidePacket(pkt, fwPacket, nb, sb, rejectBuf, i, conntrackCache.Get())
|
||||||
}
|
}
|
||||||
if err := sb.Flush(); err != nil {
|
if sb.Len() > 0 {
|
||||||
f.l.Error("Failed to write outgoing batch", "error", err, "writer", i)
|
f.flushBatch(sb, i)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Interface) flushBatch(sb batch.TxBatcher, q int) {
|
||||||
|
bufs, dsts := sb.Get()
|
||||||
|
if err := f.writers[q].WriteBatch(bufs, dsts); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||||
c.RegisterReloadCallback(f.reloadFirewall)
|
c.RegisterReloadCallback(f.reloadFirewall)
|
||||||
c.RegisterReloadCallback(f.reloadSendRecvError)
|
c.RegisterReloadCallback(f.reloadSendRecvError)
|
||||||
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
c.RegisterReloadCallback(f.reloadAcceptRecvError)
|
||||||
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
c.RegisterReloadCallback(f.reloadDisconnectInvalid)
|
||||||
c.RegisterReloadCallback(f.reloadMisc)
|
c.RegisterReloadCallback(f.reloadMisc)
|
||||||
c.RegisterReloadCallback(f.reloadEcn)
|
|
||||||
|
|
||||||
for _, udpConn := range f.writers {
|
for _, udpConn := range f.writers {
|
||||||
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
c.RegisterReloadCallback(udpConn.ReloadConfig)
|
||||||
@@ -550,20 +515,6 @@ func (f *Interface) reloadMisc(c *config.C) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// reloadEcn syncs Interface.ecnEnabled with the tunnels.ecn config knob.
|
|
||||||
// Default is enabled (RFC 6040 normal mode); set false on the rare path
|
|
||||||
// where an underlay middlebox rewrites or drops ECN bits unpredictably.
|
|
||||||
func (f *Interface) reloadEcn(c *config.C) {
|
|
||||||
initial := c.InitialLoad()
|
|
||||||
if initial || c.HasChanged("tunnels.ecn") {
|
|
||||||
v := c.GetBool("tunnels.ecn", true)
|
|
||||||
f.ecnEnabled.Store(v)
|
|
||||||
if !initial {
|
|
||||||
f.l.Info("tunnels.ecn changed", "enabled", v)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
func (f *Interface) emitStats(ctx context.Context, i time.Duration) {
|
||||||
ticker := time.NewTicker(i)
|
ticker := time.NewTicker(i)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
|
|||||||
+46
-6
@@ -15,6 +15,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
@@ -34,6 +35,7 @@ type LightHouse struct {
|
|||||||
|
|
||||||
myVpnNetworks []netip.Prefix
|
myVpnNetworks []netip.Prefix
|
||||||
myVpnNetworksTable *bart.Lite
|
myVpnNetworksTable *bart.Lite
|
||||||
|
punchConn udp.Conn
|
||||||
punchy *Punchy
|
punchy *Punchy
|
||||||
|
|
||||||
// Local cache of answers from light houses
|
// Local cache of answers from light houses
|
||||||
@@ -73,8 +75,9 @@ type LightHouse struct {
|
|||||||
|
|
||||||
calculatedRemotes atomic.Pointer[bart.Table[[]*calculatedRemote]] // Maps VpnAddr to []*calculatedRemote
|
calculatedRemotes atomic.Pointer[bart.Table[[]*calculatedRemote]] // Maps VpnAddr to []*calculatedRemote
|
||||||
|
|
||||||
metrics *MessageMetrics
|
metrics *MessageMetrics
|
||||||
l *slog.Logger
|
metricHolepunchTx metrics.Counter
|
||||||
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewLightHouseFromConfig will build a Lighthouse struct from the values provided in the config object
|
// NewLightHouseFromConfig will build a Lighthouse struct from the values provided in the config object
|
||||||
@@ -102,6 +105,7 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
addrMap: make(map[netip.Addr]*RemoteList),
|
addrMap: make(map[netip.Addr]*RemoteList),
|
||||||
nebulaPort: nebulaPort,
|
nebulaPort: nebulaPort,
|
||||||
|
punchConn: pc,
|
||||||
punchy: p,
|
punchy: p,
|
||||||
updateTrigger: make(chan struct{}, 1),
|
updateTrigger: make(chan struct{}, 1),
|
||||||
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
queryChan: make(chan netip.Addr, c.GetUint32("handshakes.query_buffer", 64)),
|
||||||
@@ -114,6 +118,9 @@ func NewLightHouseFromConfig(ctx context.Context, l *slog.Logger, c *config.C, c
|
|||||||
|
|
||||||
if c.GetBool("stats.lighthouse_metrics", false) {
|
if c.GetBool("stats.lighthouse_metrics", false) {
|
||||||
h.metrics = newLighthouseMetrics()
|
h.metrics = newLighthouseMetrics()
|
||||||
|
h.metricHolepunchTx = metrics.GetOrRegisterCounter("messages.tx.holepunch", nil)
|
||||||
|
} else {
|
||||||
|
h.metricHolepunchTx = metrics.NilCounter{}
|
||||||
}
|
}
|
||||||
|
|
||||||
err := h.reload(c, true)
|
err := h.reload(c, true)
|
||||||
@@ -1399,25 +1406,58 @@ func (lhh *LightHouseHandler) handleHostPunchNotification(n *NebulaMeta, fromVpn
|
|||||||
return
|
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.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
lhh.l.Debug("Punching",
|
||||||
|
"vpnPeer", vpnPeer,
|
||||||
|
"logVpnAddr", logVpnAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
remoteAllowList := lhh.lh.GetRemoteAllowList()
|
remoteAllowList := lhh.lh.GetRemoteAllowList()
|
||||||
for _, a := range n.Details.V4AddrPorts {
|
for _, a := range n.Details.V4AddrPorts {
|
||||||
b := protoV4AddrPortToNetAddrPort(a)
|
b := protoV4AddrPortToNetAddrPort(a)
|
||||||
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
||||||
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
punch(b, detailsVpnAddr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, a := range n.Details.V6AddrPorts {
|
for _, a := range n.Details.V6AddrPorts {
|
||||||
b := protoV6AddrPortToNetAddrPort(a)
|
b := protoV6AddrPortToNetAddrPort(a)
|
||||||
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
if remoteAllowList.Allow(detailsVpnAddr, b.Addr()) {
|
||||||
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
punch(b, detailsVpnAddr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// This sends a nebula test packet to the host trying to contact us. In the case
|
// 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
|
// of a double nat or other difficult scenario, this may help establish
|
||||||
// a tunnel. ScheduleRespond is a no-op when punchy.respond is disabled.
|
// a tunnel.
|
||||||
lhh.lh.punchy.ScheduleRespond(detailsVpnAddr)
|
if lhh.lh.punchy.GetRespond() {
|
||||||
|
go func() {
|
||||||
|
time.Sleep(lhh.lh.punchy.GetRespondDelay())
|
||||||
|
if lhh.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
lhh.l.Debug("Sending a nebula test packet",
|
||||||
|
"vpnAddr", 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))
|
||||||
|
}()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func protoAddrToNetAddr(addr *Addr) netip.Addr {
|
func protoAddrToNetAddr(addr *Addr) netip.Addr {
|
||||||
|
|||||||
@@ -5,10 +5,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
|
||||||
_ "net/http/pprof"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"runtime"
|
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
@@ -36,9 +33,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
buildVersion = moduleVersion()
|
buildVersion = moduleVersion()
|
||||||
}
|
}
|
||||||
|
|
||||||
//todo no merge
|
|
||||||
go http.ListenAndServe(":6060", nil)
|
|
||||||
|
|
||||||
// Print the config if in test, the exit comes later
|
// Print the config if in test, the exit comes later
|
||||||
if configTest {
|
if configTest {
|
||||||
b, err := yaml.Marshal(c.Settings)
|
b, err := yaml.Marshal(c.Settings)
|
||||||
@@ -61,7 +55,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
l.Info("Firewall started", "firewallHashes", fw.GetRuleHashes())
|
||||||
|
|
||||||
ssh, err := sshd.NewSSHServer(ctx, l.With("subsystem", "sshd"))
|
ssh, err := sshd.NewSSHServer(l.With("subsystem", "sshd"))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
return nil, util.ContextualizeIfNeeded("Error while creating SSH server", err)
|
||||||
}
|
}
|
||||||
@@ -176,7 +170,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostMap := NewHostMapFromConfig(l, c)
|
hostMap := NewHostMapFromConfig(l, c)
|
||||||
punchy := NewPunchyFromConfig(l, c, udpConns[0])
|
punchy := NewPunchyFromConfig(l, c)
|
||||||
connManager := newConnectionManagerFromConfig(l, c, hostMap, punchy)
|
connManager := newConnectionManagerFromConfig(l, c, hostMap, punchy)
|
||||||
lightHouse, err := NewLightHouseFromConfig(ctx, l, c, pki.getCertState(), udpConns[0], punchy)
|
lightHouse, err := NewLightHouseFromConfig(ctx, l, c, pki.getCertState(), udpConns[0], punchy)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -190,10 +184,14 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
messageMetrics = newMessageMetricsOnlyRecvError()
|
messageMetrics = newMessageMetricsOnlyRecvError()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
useRelays := c.GetBool("relay.use_relays", DefaultUseRelays) && !c.GetBool("relay.am_relay", false)
|
||||||
|
|
||||||
handshakeConfig := HandshakeConfig{
|
handshakeConfig := HandshakeConfig{
|
||||||
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
tryInterval: c.GetDuration("handshakes.try_interval", DefaultHandshakeTryInterval),
|
||||||
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
retries: int64(c.GetInt("handshakes.retries", DefaultHandshakeRetries)),
|
||||||
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
triggerBuffer: c.GetInt("handshakes.trigger_buffer", DefaultHandshakeTriggerBuffer),
|
||||||
|
useRelays: useRelays,
|
||||||
|
|
||||||
messageMetrics: messageMetrics,
|
messageMetrics: messageMetrics,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -226,7 +224,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
punchy: punchy,
|
punchy: punchy,
|
||||||
ConntrackCacheTimeout: conntrackCacheTimeout,
|
ConntrackCacheTimeout: conntrackCacheTimeout,
|
||||||
CpuAffinity: parseCpuAffinity(c, l, routines),
|
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -244,12 +241,9 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
ifce.reloadDisconnectInvalid(c)
|
ifce.reloadDisconnectInvalid(c)
|
||||||
ifce.reloadSendRecvError(c)
|
ifce.reloadSendRecvError(c)
|
||||||
ifce.reloadAcceptRecvError(c)
|
ifce.reloadAcceptRecvError(c)
|
||||||
ifce.reloadEcn(c)
|
|
||||||
|
|
||||||
handshakeManager.f = ifce
|
handshakeManager.f = ifce
|
||||||
go handshakeManager.Run(ctx)
|
go handshakeManager.Run(ctx)
|
||||||
|
|
||||||
punchy.Start(ctx, ifce, hostMap, lightHouse)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
stats, err := newStatsServerFromConfig(ctx, l, c, buildVersion, configTest)
|
||||||
@@ -279,53 +273,6 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
}, nil
|
}, 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 {
|
func moduleVersion() string {
|
||||||
info, ok := debug.ReadBuildInfo()
|
info, ok := debug.ReadBuildInfo()
|
||||||
if !ok {
|
if !ok {
|
||||||
|
|||||||
@@ -13,8 +13,6 @@ type MessageMetrics struct {
|
|||||||
|
|
||||||
rxUnknown metrics.Counter
|
rxUnknown metrics.Counter
|
||||||
txUnknown metrics.Counter
|
txUnknown metrics.Counter
|
||||||
|
|
||||||
rxInvalid metrics.Counter
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *MessageMetrics) Rx(t header.MessageType, s header.MessageSubType, i int64) {
|
func (m *MessageMetrics) Rx(t header.MessageType, s header.MessageSubType, i int64) {
|
||||||
@@ -35,11 +33,6 @@ 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 {
|
func newMessageMetrics() *MessageMetrics {
|
||||||
gen := func(t string) [][]metrics.Counter {
|
gen := func(t string) [][]metrics.Counter {
|
||||||
@@ -63,7 +56,6 @@ func newMessageMetrics() *MessageMetrics {
|
|||||||
|
|
||||||
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
rxUnknown: metrics.GetOrRegisterCounter("messages.rx.other", nil),
|
||||||
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
txUnknown: metrics.GetOrRegisterCounter("messages.tx.other", nil),
|
||||||
rxInvalid: metrics.GetOrRegisterCounter("messages.rx.invalid", nil),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+632
-45
@@ -124,7 +124,7 @@ func (x NebulaControl_MessageType) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
func (NebulaControl_MessageType) EnumDescriptor() ([]byte, []int) {
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{6, 0}
|
return fileDescriptor_2d65afa7693df5ef, []int{8, 0}
|
||||||
}
|
}
|
||||||
|
|
||||||
type NebulaMeta struct {
|
type NebulaMeta struct {
|
||||||
@@ -489,6 +489,142 @@ func (m *NebulaPing) GetTime() uint64 {
|
|||||||
return 0
|
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 struct {
|
||||||
Type NebulaControl_MessageType `protobuf:"varint,1,opt,name=Type,proto3,enum=nebula.NebulaControl_MessageType" json:"Type,omitempty"`
|
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"`
|
InitiatorRelayIndex uint32 `protobuf:"varint,2,opt,name=InitiatorRelayIndex,proto3" json:"InitiatorRelayIndex,omitempty"`
|
||||||
@@ -503,7 +639,7 @@ func (m *NebulaControl) Reset() { *m = NebulaControl{} }
|
|||||||
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
func (m *NebulaControl) String() string { return proto.CompactTextString(m) }
|
||||||
func (*NebulaControl) ProtoMessage() {}
|
func (*NebulaControl) ProtoMessage() {}
|
||||||
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
func (*NebulaControl) Descriptor() ([]byte, []int) {
|
||||||
return fileDescriptor_2d65afa7693df5ef, []int{6}
|
return fileDescriptor_2d65afa7693df5ef, []int{8}
|
||||||
}
|
}
|
||||||
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
func (m *NebulaControl) XXX_Unmarshal(b []byte) error {
|
||||||
return m.Unmarshal(b)
|
return m.Unmarshal(b)
|
||||||
@@ -593,55 +729,65 @@ func init() {
|
|||||||
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
proto.RegisterType((*V4AddrPort)(nil), "nebula.V4AddrPort")
|
||||||
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
proto.RegisterType((*V6AddrPort)(nil), "nebula.V6AddrPort")
|
||||||
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
proto.RegisterType((*NebulaPing)(nil), "nebula.NebulaPing")
|
||||||
|
proto.RegisterType((*NebulaHandshake)(nil), "nebula.NebulaHandshake")
|
||||||
|
proto.RegisterType((*NebulaHandshakeDetails)(nil), "nebula.NebulaHandshakeDetails")
|
||||||
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
proto.RegisterType((*NebulaControl)(nil), "nebula.NebulaControl")
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
func init() { proto.RegisterFile("nebula.proto", fileDescriptor_2d65afa7693df5ef) }
|
||||||
|
|
||||||
var fileDescriptor_2d65afa7693df5ef = []byte{
|
var fileDescriptor_2d65afa7693df5ef = []byte{
|
||||||
// 665 bytes of a gzipped FileDescriptorProto
|
// 785 bytes of a gzipped FileDescriptorProto
|
||||||
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x54, 0xcd, 0x6e, 0xd3, 0x5c,
|
0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0x84, 0x55, 0xcd, 0x6e, 0xeb, 0x44,
|
||||||
0x10, 0x8d, 0x1d, 0x27, 0x69, 0x27, 0x4d, 0x3e, 0x7f, 0x53, 0x51, 0x12, 0x24, 0xac, 0xe0, 0x45,
|
0x14, 0x8e, 0x1d, 0x27, 0x4e, 0x4f, 0x7e, 0xae, 0x39, 0x15, 0xc1, 0x41, 0x22, 0x0a, 0x5e, 0x54,
|
||||||
0x55, 0xb1, 0x48, 0x51, 0x5a, 0xba, 0xa6, 0x2d, 0x42, 0xa9, 0xd4, 0x9f, 0x70, 0x55, 0x8a, 0xc4,
|
0x57, 0x2c, 0x72, 0x51, 0x5a, 0xae, 0x58, 0x72, 0x1b, 0x84, 0xd2, 0xaa, 0x3f, 0x61, 0x54, 0x8a,
|
||||||
0xce, 0xb5, 0x2f, 0x8d, 0x55, 0xc7, 0x37, 0xb5, 0x6f, 0x50, 0xf3, 0x16, 0x3c, 0x0c, 0x0f, 0x01,
|
0xc4, 0x06, 0xb9, 0xf6, 0xd0, 0x58, 0x71, 0x3c, 0xa9, 0x3d, 0x41, 0xcd, 0x5b, 0xf0, 0x30, 0x3c,
|
||||||
0xbb, 0x2e, 0x59, 0xa2, 0x66, 0xc9, 0x92, 0x17, 0x40, 0xf7, 0xfa, 0xbf, 0x31, 0xb0, 0xbb, 0x33,
|
0x04, 0xec, 0xba, 0x42, 0x2c, 0x51, 0xbb, 0x64, 0xc9, 0x0b, 0xa0, 0x19, 0xff, 0x27, 0x86, 0xbb,
|
||||||
0xe7, 0x9c, 0x99, 0xc9, 0xc9, 0x8c, 0x61, 0xcd, 0xa7, 0x97, 0x33, 0xcf, 0xea, 0x4f, 0x03, 0xc6,
|
0x9b, 0x73, 0xbe, 0xef, 0x3b, 0x73, 0xe6, 0xf3, 0x9c, 0x31, 0x74, 0x02, 0x7a, 0xb7, 0xf1, 0xed,
|
||||||
0x19, 0xd6, 0xa3, 0xc8, 0xfc, 0xa9, 0x02, 0x9c, 0xca, 0xe7, 0x09, 0xe5, 0x16, 0x0e, 0x40, 0x3b,
|
0xf1, 0x3a, 0x64, 0x9c, 0x61, 0x33, 0x8e, 0xac, 0xbf, 0x55, 0x80, 0x2b, 0xb9, 0xbc, 0xa4, 0xdc,
|
||||||
0x9f, 0x4f, 0x69, 0x47, 0xe9, 0x29, 0x5b, 0xed, 0x81, 0xd1, 0x8f, 0x35, 0x19, 0xa3, 0x7f, 0x42,
|
0xc6, 0x09, 0x68, 0x37, 0xdb, 0x35, 0x35, 0x95, 0x91, 0xf2, 0xba, 0x37, 0x19, 0x8e, 0x13, 0x4d,
|
||||||
0xc3, 0xd0, 0xba, 0xa2, 0x82, 0x45, 0x24, 0x17, 0x77, 0xa0, 0xf1, 0x9a, 0x72, 0xcb, 0xf5, 0xc2,
|
0xce, 0x18, 0x5f, 0xd2, 0x28, 0xb2, 0xef, 0xa9, 0x60, 0x11, 0xc9, 0xc5, 0x63, 0xd0, 0xbf, 0xa6,
|
||||||
0x8e, 0xda, 0x53, 0xb6, 0x9a, 0x83, 0xee, 0xb2, 0x2c, 0x26, 0x90, 0x84, 0x69, 0xfe, 0x52, 0xa0,
|
0xdc, 0xf6, 0xfc, 0xc8, 0x54, 0x47, 0xca, 0xeb, 0xf6, 0x64, 0xb0, 0x2f, 0x4b, 0x08, 0x24, 0x65,
|
||||||
0x99, 0x2b, 0x85, 0x2b, 0xa0, 0x9d, 0x32, 0x9f, 0xea, 0x15, 0x6c, 0xc1, 0xea, 0x90, 0x85, 0xfc,
|
0x5a, 0xff, 0x28, 0xd0, 0x2e, 0x94, 0xc2, 0x16, 0x68, 0x57, 0x2c, 0xa0, 0x46, 0x0d, 0xbb, 0x70,
|
||||||
0xed, 0x8c, 0x06, 0x73, 0x5d, 0x41, 0x84, 0x76, 0x1a, 0x12, 0x3a, 0xf5, 0xe6, 0xba, 0x8a, 0x4f,
|
0x30, 0x63, 0x11, 0xff, 0x76, 0x43, 0xc3, 0xad, 0xa1, 0x20, 0x42, 0x2f, 0x0b, 0x09, 0x5d, 0xfb,
|
||||||
0x60, 0x43, 0xe4, 0xde, 0x4d, 0x1d, 0x8b, 0xd3, 0x53, 0xc6, 0xdd, 0x8f, 0xae, 0x6d, 0x71, 0x97,
|
0x5b, 0x43, 0xc5, 0x8f, 0xa1, 0x2f, 0x72, 0xdf, 0xad, 0x5d, 0x9b, 0xd3, 0x2b, 0xc6, 0xbd, 0x9f,
|
||||||
0xf9, 0x7a, 0x15, 0xbb, 0xf0, 0x48, 0x60, 0x27, 0xec, 0x13, 0x75, 0x0a, 0x90, 0x96, 0x40, 0xa3,
|
0x3c, 0xc7, 0xe6, 0x1e, 0x0b, 0x8c, 0x3a, 0x0e, 0xe0, 0x43, 0x81, 0x5d, 0xb2, 0x9f, 0xa9, 0x5b,
|
||||||
0x99, 0x6f, 0x8f, 0x0b, 0x50, 0x0d, 0xdb, 0x00, 0x02, 0x7a, 0x3f, 0x66, 0xd6, 0xc4, 0xd5, 0xeb,
|
0x82, 0xb4, 0x14, 0x9a, 0x6f, 0x02, 0x67, 0x51, 0x82, 0x1a, 0xd8, 0x03, 0x10, 0xd0, 0xf7, 0x0b,
|
||||||
0xb8, 0x0e, 0xff, 0x65, 0x71, 0xd4, 0xb6, 0x21, 0x26, 0x1b, 0x59, 0x7c, 0x7c, 0x38, 0xa6, 0xf6,
|
0x66, 0xaf, 0x3c, 0xa3, 0x89, 0x87, 0xf0, 0x2a, 0x8f, 0xe3, 0x6d, 0x75, 0xd1, 0xd9, 0xdc, 0xe6,
|
||||||
0xb5, 0xbe, 0x22, 0x26, 0x4b, 0xc3, 0x88, 0xb2, 0x8a, 0x4f, 0xa1, 0x5b, 0x3e, 0xd9, 0xbe, 0x7d,
|
0x8b, 0xe9, 0x82, 0x3a, 0x4b, 0xa3, 0x25, 0x3a, 0xcb, 0xc2, 0x98, 0x72, 0x80, 0x9f, 0xc0, 0xa0,
|
||||||
0xad, 0x83, 0xf9, 0x4d, 0x85, 0xff, 0x97, 0x4c, 0x41, 0x13, 0xe0, 0xcc, 0x73, 0x2e, 0xa6, 0xfe,
|
0xba, 0xb3, 0x77, 0xce, 0xd2, 0x00, 0xeb, 0x77, 0x15, 0x3e, 0xd8, 0x33, 0x05, 0x2d, 0x80, 0x6b,
|
||||||
0xbe, 0xe3, 0x04, 0xd2, 0xfa, 0xd6, 0x81, 0xda, 0x51, 0x48, 0x2e, 0x8b, 0x9b, 0xd0, 0x48, 0x08,
|
0xdf, 0xbd, 0x5d, 0x07, 0xef, 0x5c, 0x37, 0x94, 0xd6, 0x77, 0x4f, 0x55, 0x53, 0x21, 0x85, 0x2c,
|
||||||
0x75, 0x69, 0xf2, 0x5a, 0x62, 0xb2, 0xc8, 0x91, 0x04, 0xc4, 0x3e, 0xe8, 0x67, 0x9e, 0x43, 0xa8,
|
0x1e, 0x81, 0x9e, 0x12, 0x9a, 0xd2, 0xe4, 0x4e, 0x6a, 0xb2, 0xc8, 0x91, 0x14, 0xc4, 0x31, 0x18,
|
||||||
0x67, 0xcd, 0xe3, 0x54, 0xd8, 0xa9, 0xf5, 0xaa, 0x71, 0xc5, 0x25, 0x0c, 0x07, 0xd0, 0x2a, 0x92,
|
0xd7, 0xbe, 0x4b, 0xa8, 0x6f, 0x6f, 0x93, 0x54, 0x64, 0x36, 0x46, 0xf5, 0xa4, 0xe2, 0x1e, 0x86,
|
||||||
0x1b, 0xbd, 0xea, 0x52, 0xf5, 0x22, 0x05, 0x77, 0xa1, 0x79, 0xb1, 0x2b, 0x9e, 0x23, 0x16, 0x70,
|
0x13, 0xe8, 0x96, 0xc9, 0xfa, 0xa8, 0xbe, 0x57, 0xbd, 0x4c, 0xc1, 0x13, 0x68, 0xdf, 0x9e, 0x88,
|
||||||
0xf1, 0xa7, 0x0b, 0x05, 0x26, 0x8a, 0x0c, 0x22, 0x79, 0x9a, 0x54, 0xed, 0x65, 0x2a, 0xed, 0x81,
|
0xe5, 0x9c, 0x85, 0x5c, 0x7c, 0x74, 0xa1, 0xc0, 0x54, 0x91, 0x43, 0xa4, 0x48, 0x93, 0xaa, 0xb7,
|
||||||
0x6a, 0x2f, 0xa7, 0xca, 0x68, 0xd8, 0x81, 0x86, 0xcd, 0x66, 0x3e, 0xa7, 0x41, 0xa7, 0x2a, 0x8c,
|
0xb9, 0x4a, 0xdb, 0x51, 0xbd, 0x2d, 0xa8, 0x72, 0x1a, 0x9a, 0xa0, 0x3b, 0x6c, 0x13, 0x70, 0x1a,
|
||||||
0x21, 0x49, 0x68, 0x6e, 0x82, 0x26, 0x7f, 0x71, 0x1b, 0xd4, 0xa1, 0x2b, 0x5d, 0xd3, 0x88, 0x3a,
|
0x9a, 0x75, 0x61, 0x0c, 0x49, 0x43, 0xeb, 0x08, 0x34, 0x79, 0xe2, 0x1e, 0xa8, 0x33, 0x4f, 0xba,
|
||||||
0x74, 0x45, 0x7c, 0xcc, 0xe4, 0x26, 0x6a, 0x44, 0x3d, 0x66, 0xe6, 0x2e, 0x40, 0x36, 0x06, 0x62,
|
0xa6, 0x11, 0x75, 0xe6, 0x89, 0xf8, 0x82, 0xc9, 0x9b, 0xa8, 0x11, 0xf5, 0x82, 0x59, 0x27, 0x00,
|
||||||
0xa4, 0x8a, 0x5c, 0x26, 0x51, 0x05, 0x04, 0x4d, 0x60, 0x52, 0xd3, 0x22, 0xf2, 0x6d, 0xbe, 0x02,
|
0x79, 0x1b, 0x88, 0xb1, 0x2a, 0x76, 0x99, 0xc4, 0x15, 0x10, 0x34, 0x81, 0x49, 0x4d, 0x97, 0xc8,
|
||||||
0xc8, 0xc6, 0xf8, 0x57, 0x8f, 0xb4, 0x42, 0x35, 0x57, 0xe1, 0x36, 0x39, 0xac, 0x91, 0xeb, 0x5f,
|
0xb5, 0xf5, 0x15, 0x40, 0xde, 0xc6, 0xfb, 0xf6, 0xc8, 0x2a, 0xd4, 0x0b, 0x15, 0x1e, 0xd3, 0xc1,
|
||||||
0xfd, 0xfd, 0xb0, 0x04, 0xa3, 0xe4, 0xb0, 0x10, 0xb4, 0x73, 0x77, 0x42, 0xe3, 0x3e, 0xf2, 0x6d,
|
0x9a, 0x7b, 0xc1, 0xfd, 0xff, 0x0f, 0x96, 0x60, 0x54, 0x0c, 0x16, 0x82, 0x76, 0xe3, 0xad, 0x68,
|
||||||
0x9a, 0x4b, 0x67, 0x23, 0xc4, 0x7a, 0x05, 0x57, 0xa1, 0x16, 0x2d, 0xa1, 0x62, 0x7e, 0xa9, 0x42,
|
0xb2, 0x8f, 0x5c, 0x5b, 0xd6, 0xde, 0xd8, 0x08, 0xb1, 0x51, 0xc3, 0x03, 0x68, 0xc4, 0x97, 0x50,
|
||||||
0x2b, 0x2a, 0x7c, 0xc8, 0x7c, 0x1e, 0x30, 0x0f, 0x5f, 0x16, 0xba, 0x3f, 0x2b, 0x76, 0x8f, 0x49,
|
0xb1, 0x7e, 0x84, 0x57, 0x71, 0xdd, 0x99, 0x1d, 0xb8, 0xd1, 0xc2, 0x5e, 0x52, 0xfc, 0x32, 0x9f,
|
||||||
0x25, 0x03, 0xbc, 0x80, 0xf5, 0x23, 0xdf, 0xe5, 0xae, 0xc5, 0x59, 0x20, 0x57, 0xe0, 0xc8, 0x77,
|
0x51, 0x45, 0x5e, 0x9f, 0x9d, 0x0e, 0x32, 0xe6, 0xee, 0xa0, 0x8a, 0x26, 0x66, 0x2b, 0xdb, 0x91,
|
||||||
0xe8, 0x6d, 0xec, 0x53, 0x19, 0x24, 0x14, 0x84, 0x86, 0x53, 0xe6, 0x3b, 0x34, 0xaf, 0x88, 0x7c,
|
0x4d, 0x74, 0x88, 0x5c, 0x5b, 0x7f, 0x28, 0xd0, 0xaf, 0xd6, 0x09, 0xfa, 0x94, 0x86, 0x5c, 0xee,
|
||||||
0x29, 0x83, 0xf0, 0x39, 0xb4, 0x93, 0xa5, 0x3c, 0x67, 0xf2, 0xaf, 0xd1, 0xd2, 0x03, 0x78, 0x80,
|
0xd2, 0x21, 0x72, 0x8d, 0x47, 0xd0, 0x3b, 0x0b, 0x3c, 0xee, 0xd9, 0x9c, 0x85, 0x67, 0x81, 0x4b,
|
||||||
0xe4, 0x97, 0xfb, 0x4d, 0xc0, 0x26, 0x92, 0x5d, 0x4b, 0xd9, 0x4b, 0x18, 0xf6, 0xa1, 0x99, 0x2f,
|
0x1f, 0x13, 0xa7, 0x77, 0xb2, 0x82, 0x47, 0x68, 0xb4, 0x66, 0x81, 0x4b, 0x13, 0x5e, 0xec, 0xe7,
|
||||||
0x5c, 0x76, 0x38, 0x79, 0x42, 0x7a, 0x0c, 0x69, 0xf1, 0x46, 0x89, 0xa2, 0x48, 0x31, 0x87, 0x7f,
|
0x4e, 0x16, 0xfb, 0xd0, 0x9c, 0x32, 0xb6, 0xf4, 0xa8, 0xa9, 0x49, 0x67, 0x92, 0x28, 0xf3, 0xab,
|
||||||
0xfa, 0x8e, 0x6d, 0x00, 0x1e, 0x06, 0xd4, 0xe2, 0x54, 0xf2, 0x09, 0xbd, 0x99, 0xd1, 0x90, 0xeb,
|
0x91, 0xfb, 0x85, 0x23, 0x68, 0x8b, 0x1e, 0x6e, 0x69, 0x18, 0x79, 0x2c, 0x30, 0x5b, 0xb2, 0x60,
|
||||||
0x0a, 0x3e, 0x86, 0xf5, 0x42, 0x5e, 0x58, 0x12, 0x52, 0x5d, 0x3d, 0xd8, 0xf9, 0x7a, 0x6f, 0x28,
|
0x31, 0x75, 0xae, 0xb5, 0x9a, 0x86, 0x7e, 0xae, 0xb5, 0x74, 0xa3, 0x65, 0xfd, 0x5a, 0x87, 0x6e,
|
||||||
0x77, 0xf7, 0x86, 0xf2, 0xe3, 0xde, 0x50, 0x3e, 0x2f, 0x8c, 0xca, 0xdd, 0xc2, 0xa8, 0x7c, 0x5f,
|
0x7c, 0xb0, 0x29, 0x0b, 0x78, 0xc8, 0x7c, 0xfc, 0xa2, 0xf4, 0xdd, 0x3e, 0x2d, 0xbb, 0x96, 0x90,
|
||||||
0x18, 0x95, 0x0f, 0xdd, 0x2b, 0x97, 0x8f, 0x67, 0x97, 0x7d, 0x9b, 0x4d, 0xb6, 0x43, 0xcf, 0xb2,
|
0x2a, 0x3e, 0xdd, 0xe7, 0x70, 0x98, 0x1d, 0x4e, 0x0e, 0x4f, 0xf1, 0xdc, 0x55, 0x90, 0x50, 0x64,
|
||||||
0xaf, 0xc7, 0x37, 0xdb, 0xd1, 0x48, 0x97, 0x75, 0xf9, 0x39, 0xdf, 0xf9, 0x1d, 0x00, 0x00, 0xff,
|
0xc7, 0x2c, 0x28, 0x62, 0x07, 0xaa, 0x20, 0xfc, 0x0c, 0x7a, 0xe9, 0x38, 0xdf, 0x30, 0x79, 0xa9,
|
||||||
0xff, 0x51, 0x0a, 0xe3, 0xd7, 0xde, 0x05, 0x00, 0x00,
|
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,
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
func (m *NebulaMeta) Marshal() (dAtA []byte, err error) {
|
||||||
@@ -926,6 +1072,103 @@ func (m *NebulaPing) MarshalToSizedBuffer(dAtA []byte) (int, error) {
|
|||||||
return len(dAtA) - i, nil
|
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) {
|
func (m *NebulaControl) Marshal() (dAtA []byte, err error) {
|
||||||
size := m.Size()
|
size := m.Size()
|
||||||
dAtA = make([]byte, size)
|
dAtA = make([]byte, size)
|
||||||
@@ -1132,6 +1375,51 @@ func (m *NebulaPing) Size() (n int) {
|
|||||||
return n
|
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) {
|
func (m *NebulaControl) Size() (n int) {
|
||||||
if m == nil {
|
if m == nil {
|
||||||
return 0
|
return 0
|
||||||
@@ -1948,6 +2236,305 @@ func (m *NebulaPing) Unmarshal(dAtA []byte) error {
|
|||||||
}
|
}
|
||||||
return nil
|
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 {
|
func (m *NebulaControl) Unmarshal(dAtA []byte) error {
|
||||||
l := len(dAtA)
|
l := len(dAtA)
|
||||||
iNdEx := 0
|
iNdEx := 0
|
||||||
|
|||||||
+15
-3
@@ -60,9 +60,21 @@ message NebulaPing {
|
|||||||
uint64 Time = 2;
|
uint64 Time = 2;
|
||||||
}
|
}
|
||||||
|
|
||||||
// NebulaHandshake / NebulaHandshakeDetails moved to
|
message NebulaHandshake {
|
||||||
// handshake/handshake.proto. The handshake package speaks that wire format
|
NebulaHandshakeDetails Details = 1;
|
||||||
// directly via a hand-written encoder/decoder.
|
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;
|
||||||
|
}
|
||||||
|
|
||||||
message NebulaControl {
|
message NebulaControl {
|
||||||
enum MessageType {
|
enum MessageType {
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
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 cipher.AEAD
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewNebulaCipherState(s *noise.CipherState) *NebulaCipherState {
|
||||||
|
x := s.Cipher()
|
||||||
|
return &NebulaCipherState{c: x.(cipher.AEAD)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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.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.Open(out, nb, ciphertext, ad)
|
||||||
|
} else {
|
||||||
|
return []byte{}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *NebulaCipherState) Overhead() int {
|
||||||
|
if s != nil {
|
||||||
|
return s.c.Overhead()
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
@@ -1,53 +0,0 @@
|
|||||||
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()
|
|
||||||
}
|
|
||||||
@@ -1,52 +0,0 @@
|
|||||||
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()
|
|
||||||
}
|
|
||||||
@@ -1,40 +0,0 @@
|
|||||||
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()))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,166 +0,0 @@
|
|||||||
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())
|
|
||||||
}
|
|
||||||
+411
-253
@@ -2,57 +2,41 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/gopacket/layers"
|
||||||
|
"golang.org/x/net/ipv6"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay/batch"
|
"golang.org/x/net/ipv4"
|
||||||
"github.com/slackhq/nebula/udp"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
var ErrOutOfWindow = errors.New("out of window packet")
|
const (
|
||||||
|
minFwPacketLen = 4
|
||||||
|
)
|
||||||
|
|
||||||
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) {
|
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) {
|
||||||
err := h.Parse(packet)
|
err := h.Parse(packet)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
// Hole punch packets are 0 or 1 byte big, so lets ignore printing those errors
|
||||||
// TODO: record metrics for rx holepunch/punchy packets?
|
|
||||||
if len(packet) > 1 {
|
if len(packet) > 1 {
|
||||||
f.messageMetrics.RxInvalid(1)
|
f.l.Info("Error while parsing inbound packet",
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
"from", via,
|
||||||
f.l.Debug("Error while parsing inbound packet",
|
"error", err,
|
||||||
"from", via,
|
"packet", packet,
|
||||||
"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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
//l.Error("in packet ", header, packet[HeaderLen:])
|
||||||
if !via.IsRelayed {
|
if !via.IsRelayed {
|
||||||
if f.myVpnNetworksTable.Contains(via.UdpAddr.Addr()) {
|
if f.myVpnNetworksTable.Contains(via.UdpAddr.Addr()) {
|
||||||
f.messageMetrics.RxInvalid(1)
|
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("Refusing to process double encrypted packet", "from", via)
|
f.l.Debug("Refusing to process double encrypted packet", "from", via)
|
||||||
}
|
}
|
||||||
@@ -60,193 +44,215 @@ func (f *Interface) readOutsidePackets(via ViaSender, out []byte, packet []byte,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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
|
var hostinfo *HostInfo
|
||||||
if isMessageRelay {
|
// verify if we've seen this index before, otherwise respond to the handshake initiation
|
||||||
|
if h.Type == header.Message && h.Subtype == header.MessageRelay {
|
||||||
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
hostinfo = f.hostMap.QueryRelayIndex(h.RemoteIndex)
|
||||||
} else {
|
} else {
|
||||||
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
hostinfo = f.hostMap.QueryIndex(h.RemoteIndex)
|
||||||
}
|
}
|
||||||
|
|
||||||
// At this point we should have a valid existing tunnel, verify and send
|
var ci *ConnectionState
|
||||||
// recvError if necessary
|
if hostinfo != nil {
|
||||||
if hostinfo == nil || hostinfo.ConnectionState == nil {
|
ci = hostinfo.ConnectionState
|
||||||
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 {
|
switch h.Type {
|
||||||
case header.Message:
|
case header.Message:
|
||||||
|
if !f.handleEncrypted(ci, via, h) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
switch h.Subtype {
|
switch h.Subtype {
|
||||||
case header.MessageNone:
|
case header.MessageNone:
|
||||||
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, parsedRx, nb, q, localCache, meta)
|
if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) {
|
||||||
default:
|
return
|
||||||
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
}
|
||||||
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).Error("HostInfo missing remote relay index",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
"remoteIndex", h.RemoteIndex,
|
||||||
|
)
|
||||||
|
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).Info("Failed to find target host info by ip",
|
||||||
|
"relayTo", relay.PeerAddr,
|
||||||
|
"error", err,
|
||||||
|
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
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).Info("Unexpected target relay state",
|
||||||
|
"relayTo", relay.PeerAddr,
|
||||||
|
"relayFrom", hostinfo.vpnAddrs[0],
|
||||||
|
"targetRelayState", targetRelay.State,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
case header.LightHouse:
|
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).Error("Failed to decrypt lighthouse packet",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"packet", packet,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
//TODO: assert via is not relayed
|
//TODO: assert via is not relayed
|
||||||
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, d, f)
|
||||||
|
|
||||||
|
// Fallthrough to the bottom to record incoming traffic
|
||||||
|
|
||||||
case header.Test:
|
case header.Test:
|
||||||
switch h.Subtype {
|
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||||
case header.TestReply:
|
if !f.handleEncrypted(ci, via, h) {
|
||||||
// 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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to decrypt test packet",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"packet", 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:
|
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).Error("Failed to decrypt CloseTunnel packet",
|
||||||
|
"error", err,
|
||||||
|
"from", via,
|
||||||
|
"packet", packet,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
hostinfo.logger(f.l).Info("Close tunnel received, tearing down.", "from", via)
|
hostinfo.logger(f.l).Info("Close tunnel received, tearing down.", "from", via)
|
||||||
|
|
||||||
f.closeTunnel(hostinfo)
|
f.closeTunnel(hostinfo)
|
||||||
|
return
|
||||||
|
|
||||||
case header.Control:
|
case header.Control:
|
||||||
f.relayManager.HandleControlMsg(hostinfo, out, f)
|
if !f.handleEncrypted(ci, via, h) {
|
||||||
|
return
|
||||||
default:
|
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
||||||
case ForwardingType:
|
|
||||||
// Find the target HostInfo relay object
|
|
||||||
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
hostinfo.logger(f.l).Error("Failed to decrypt Control packet",
|
||||||
"relayTo", relay.PeerAddr,
|
|
||||||
"error", err,
|
"error", err,
|
||||||
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
|
"from", via,
|
||||||
|
"packet", packet,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// If that relay is Established, forward the payload through it
|
f.relayManager.HandleControlMsg(hostinfo, d, f)
|
||||||
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:
|
default:
|
||||||
|
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("Unexpected relay type", "from", via, "relayType", relay.Type)
|
hostinfo.logger(f.l).Debug("Unexpected packet received", "from", via)
|
||||||
}
|
}
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
f.handleHostRoaming(hostinfo, via)
|
||||||
|
|
||||||
|
f.connectionManager.In(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
||||||
@@ -294,6 +300,23 @@ 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 (
|
var (
|
||||||
ErrPacketTooShort = errors.New("packet is too short")
|
ErrPacketTooShort = errors.New("packet is too short")
|
||||||
ErrUnknownIPVersion = errors.New("packet is an unknown ip version")
|
ErrUnknownIPVersion = errors.New("packet is an unknown ip version")
|
||||||
@@ -304,16 +327,191 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
// newPacket validates and parses the interesting bits for the firewall out of the ip and sub protocol headers
|
||||||
// newPacket parses data into a fully-hydrated firewall.Packet — kept as a
|
|
||||||
// thin wrapper around newPacketKey + Hydrate so there's one source of
|
|
||||||
// parse logic. Callers that don't need the netip.Addr-rich form (e.g.
|
|
||||||
// conntrack-only paths) should use newPacketKey directly.
|
|
||||||
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
func newPacket(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||||
var parsed batch.RxParsed
|
if len(data) < 1 {
|
||||||
if err := batch.ParsePacket(data, incoming, &parsed); err != nil {
|
return ErrPacketTooShort
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
parsed.Key.Hydrate(fp)
|
|
||||||
|
version := int((data[0] >> 4) & 0x0f)
|
||||||
|
switch version {
|
||||||
|
case ipv4.Version:
|
||||||
|
return parseV4(data, incoming, fp)
|
||||||
|
case ipv6.Version:
|
||||||
|
return parseV6(data, incoming, fp)
|
||||||
|
}
|
||||||
|
return ErrUnknownIPVersion
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseV6(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||||
|
dataLen := len(data)
|
||||||
|
if dataLen < ipv6.HeaderLen {
|
||||||
|
return ErrIPv6PacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
if incoming {
|
||||||
|
fp.RemoteAddr, _ = netip.AddrFromSlice(data[8:24])
|
||||||
|
fp.LocalAddr, _ = netip.AddrFromSlice(data[24:40])
|
||||||
|
} else {
|
||||||
|
fp.LocalAddr, _ = netip.AddrFromSlice(data[8:24])
|
||||||
|
fp.RemoteAddr, _ = netip.AddrFromSlice(data[24:40])
|
||||||
|
}
|
||||||
|
|
||||||
|
protoAt := 6 // NextHeader is at 6 bytes into the ipv6 header
|
||||||
|
offset := ipv6.HeaderLen // Start at the end of the ipv6 header
|
||||||
|
next := 0
|
||||||
|
for {
|
||||||
|
if protoAt >= dataLen {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
proto := layers.IPProtocol(data[protoAt])
|
||||||
|
|
||||||
|
switch proto {
|
||||||
|
case layers.IPProtocolESP, layers.IPProtocolNoNextHeader:
|
||||||
|
fp.Protocol = uint8(proto)
|
||||||
|
fp.RemotePort = 0
|
||||||
|
fp.LocalPort = 0
|
||||||
|
fp.Fragment = false
|
||||||
|
return nil
|
||||||
|
|
||||||
|
case layers.IPProtocolICMPv6:
|
||||||
|
if dataLen < offset+6 {
|
||||||
|
return ErrIPv6PacketTooShort
|
||||||
|
}
|
||||||
|
fp.Protocol = uint8(proto)
|
||||||
|
fp.LocalPort = 0 //incoming vs outgoing doesn't matter for icmpv6
|
||||||
|
icmptype := data[offset+1]
|
||||||
|
switch icmptype {
|
||||||
|
case layers.ICMPv6TypeEchoRequest, layers.ICMPv6TypeEchoReply:
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset+4 : offset+6]) //identifier
|
||||||
|
default:
|
||||||
|
fp.RemotePort = 0
|
||||||
|
}
|
||||||
|
fp.Fragment = false
|
||||||
|
return nil
|
||||||
|
|
||||||
|
case layers.IPProtocolTCP, layers.IPProtocolUDP:
|
||||||
|
if dataLen < offset+4 {
|
||||||
|
return ErrIPv6PacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
fp.Protocol = uint8(proto)
|
||||||
|
if incoming {
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset : offset+2])
|
||||||
|
fp.LocalPort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
||||||
|
} else {
|
||||||
|
fp.LocalPort = binary.BigEndian.Uint16(data[offset : offset+2])
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[offset+2 : offset+4])
|
||||||
|
}
|
||||||
|
|
||||||
|
fp.Fragment = false
|
||||||
|
return nil
|
||||||
|
|
||||||
|
case layers.IPProtocolIPv6Fragment:
|
||||||
|
// Fragment header is 8 bytes, need at least offset+4 to read the offset field
|
||||||
|
if dataLen < offset+8 {
|
||||||
|
return ErrIPv6PacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is the first fragment
|
||||||
|
fragmentOffset := binary.BigEndian.Uint16(data[offset+2:offset+4]) &^ uint16(0x7) // Remove the reserved and M flag bits
|
||||||
|
if fragmentOffset != 0 {
|
||||||
|
// Non-first fragment, use what we have now and stop processing
|
||||||
|
fp.Protocol = data[offset]
|
||||||
|
fp.Fragment = true
|
||||||
|
fp.RemotePort = 0
|
||||||
|
fp.LocalPort = 0
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// The next loop should be the transport layer since we are the first fragment
|
||||||
|
next = 8 // Fragment headers are always 8 bytes
|
||||||
|
|
||||||
|
case layers.IPProtocolAH:
|
||||||
|
// Auth headers, used by IPSec, have a different meaning for header length
|
||||||
|
if dataLen <= offset+1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
next = int(data[offset+1]+2) << 2
|
||||||
|
|
||||||
|
default:
|
||||||
|
// Normal ipv6 header length processing
|
||||||
|
if dataLen <= offset+1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
next = int(data[offset+1]+1) << 3
|
||||||
|
}
|
||||||
|
|
||||||
|
if next <= 0 {
|
||||||
|
// Safety check, each ipv6 header has to be at least 8 bytes
|
||||||
|
next = 8
|
||||||
|
}
|
||||||
|
|
||||||
|
protoAt = offset
|
||||||
|
offset = offset + next
|
||||||
|
}
|
||||||
|
|
||||||
|
return ErrIPv6CouldNotFindPayload
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseV4(data []byte, incoming bool, fp *firewall.Packet) error {
|
||||||
|
// Do we at least have an ipv4 header worth of data?
|
||||||
|
if len(data) < ipv4.HeaderLen {
|
||||||
|
return ErrIPv4PacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
// Adjust our start position based on the advertised ip header length
|
||||||
|
ihl := int(data[0]&0x0f) << 2
|
||||||
|
|
||||||
|
// Well-formed ip header length?
|
||||||
|
if ihl < ipv4.HeaderLen {
|
||||||
|
return ErrIPv4InvalidHeaderLength
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check if this is the second or further fragment of a fragmented packet.
|
||||||
|
flagsfrags := binary.BigEndian.Uint16(data[6:8])
|
||||||
|
fp.Fragment = (flagsfrags & 0x1FFF) != 0
|
||||||
|
|
||||||
|
// Firewall handles protocol checks
|
||||||
|
fp.Protocol = data[9]
|
||||||
|
|
||||||
|
// Accounting for a variable header length, do we have enough data for our src/dst tuples?
|
||||||
|
minLen := ihl
|
||||||
|
if !fp.Fragment {
|
||||||
|
if fp.Protocol == firewall.ProtoICMP {
|
||||||
|
minLen += minFwPacketLen + 2
|
||||||
|
} else {
|
||||||
|
minLen += minFwPacketLen
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(data) < minLen {
|
||||||
|
return ErrIPv4InvalidHeaderLength
|
||||||
|
}
|
||||||
|
|
||||||
|
if incoming { // Firewall packets are locally oriented
|
||||||
|
fp.RemoteAddr, _ = netip.AddrFromSlice(data[12:16])
|
||||||
|
fp.LocalAddr, _ = netip.AddrFromSlice(data[16:20])
|
||||||
|
} else {
|
||||||
|
fp.LocalAddr, _ = netip.AddrFromSlice(data[12:16])
|
||||||
|
fp.RemoteAddr, _ = netip.AddrFromSlice(data[16:20])
|
||||||
|
}
|
||||||
|
|
||||||
|
if fp.Fragment {
|
||||||
|
fp.RemotePort = 0
|
||||||
|
fp.LocalPort = 0
|
||||||
|
} else if fp.Protocol == firewall.ProtoICMP { //note that orientation doesn't matter on ICMP
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+4 : ihl+6]) //identifier
|
||||||
|
fp.LocalPort = 0 //code would be uint16(data[ihl+1])
|
||||||
|
} else if incoming {
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[ihl : ihl+2]) //src port
|
||||||
|
fp.LocalPort = binary.BigEndian.Uint16(data[ihl+2 : ihl+4]) //dst port
|
||||||
|
} else {
|
||||||
|
fp.LocalPort = binary.BigEndian.Uint16(data[ihl : ihl+2]) //src port
|
||||||
|
fp.RemotePort = binary.BigEndian.Uint16(data[ihl+2 : ihl+4]) //dst port
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -325,83 +523,41 @@ func (f *Interface) decrypt(hostinfo *HostInfo, mc uint64, out []byte, packet []
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
if !hostinfo.ConnectionState.window.Update(f.l, mc) {
|
||||||
return nil, ErrOutOfWindow
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("dropping out of window packet", "header", h)
|
||||||
|
}
|
||||||
|
return nil, errors.New("out of window packet")
|
||||||
}
|
}
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// 2-bit IP-level ECN codepoints (lower bits of IPv4 ToS / IPv6 TC).
|
func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) bool {
|
||||||
const (
|
var err error
|
||||||
ecnNotECT = 0x00
|
|
||||||
ecnECT1 = 0x01
|
|
||||||
ecnECT0 = 0x02
|
|
||||||
ecnCE = 0x03
|
|
||||||
)
|
|
||||||
|
|
||||||
// applyOuterECN folds an outer CE mark from the underlay into the inner
|
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
||||||
// IP header per RFC 6040 normal mode. It mutates pkt[1] in place. Other
|
if err != nil {
|
||||||
// codepoints are advisory only and leave the inner unchanged.
|
hostinfo.logger(f.l).Error("Failed to decrypt packet", "error", err)
|
||||||
//
|
return false
|
||||||
// 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)
|
err = newPacket(out, true, fwPacket)
|
||||||
// and the batcher hint (parsedRx.tcp/udp). Replaces newPacket — and
|
|
||||||
// pointedly does NOT fill fwPacket.LocalAddr/RemoteAddr, since
|
|
||||||
// firewall.Drop's fast path uses Key alone and only hydrates fwPacket
|
|
||||||
// from Key on the slow path.
|
|
||||||
*fwPacket = firewall.Packet{}
|
|
||||||
err := batch.ParsePacket(out, true, parsedRx)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
hostinfo.logger(f.l).Warn("Error while validating inbound packet",
|
||||||
"error", err,
|
"error", err,
|
||||||
"packet", out,
|
"packet", out,
|
||||||
)
|
)
|
||||||
return
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
dropReason := f.firewall.Drop(parsedRx.Key, fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
if !hostinfo.ConnectionState.window.Update(f.l, messageCounter) {
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("dropping out of window packet", "fwPacket", fwPacket)
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
if dropReason != nil {
|
if dropReason != nil {
|
||||||
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
// NOTE: We give `packet` as the `out` here since we already decrypted from it and we don't need it anymore
|
||||||
// This gives us a buffer to build the reject packet in
|
// This gives us a buffer to build the reject packet in
|
||||||
@@ -412,13 +568,15 @@ func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, p
|
|||||||
"reason", dropReason,
|
"reason", dropReason,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
return
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
err = f.batchers[q].CommitInbound(out, parsedRx)
|
f.connectionManager.In(hostinfo)
|
||||||
|
err = f.batchers[q].Commit(out)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
||||||
|
|||||||
+19
-20
@@ -11,7 +11,6 @@ import (
|
|||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/overlay/batch"
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"golang.org/x/net/ipv4"
|
"golang.org/x/net/ipv4"
|
||||||
@@ -22,13 +21,13 @@ func Test_newPacket(t *testing.T) {
|
|||||||
|
|
||||||
// length fails
|
// length fails
|
||||||
err := newPacket([]byte{}, true, p)
|
err := newPacket([]byte{}, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrPacketTooShort)
|
require.ErrorIs(t, err, ErrPacketTooShort)
|
||||||
|
|
||||||
err = newPacket([]byte{0x40}, true, p)
|
err = newPacket([]byte{0x40}, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv4PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv4PacketTooShort)
|
||||||
|
|
||||||
err = newPacket([]byte{0x60}, true, p)
|
err = newPacket([]byte{0x60}, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// length fail with ip options
|
// length fail with ip options
|
||||||
h := ipv4.Header{
|
h := ipv4.Header{
|
||||||
@@ -41,15 +40,15 @@ func Test_newPacket(t *testing.T) {
|
|||||||
|
|
||||||
b, _ := h.Marshal()
|
b, _ := h.Marshal()
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv4InvalidHeaderLength)
|
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
||||||
|
|
||||||
// not an ipv4 packet
|
// not an ipv4 packet
|
||||||
err = newPacket([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p)
|
err = newPacket([]byte{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrUnknownIPVersion)
|
require.ErrorIs(t, err, ErrUnknownIPVersion)
|
||||||
|
|
||||||
// invalid ihl
|
// invalid ihl
|
||||||
err = newPacket([]byte{4<<4 | (8 >> 2 & 0x0f), 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p)
|
err = newPacket([]byte{4<<4 | (8 >> 2 & 0x0f), 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0}, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv4InvalidHeaderLength)
|
require.ErrorIs(t, err, ErrIPv4InvalidHeaderLength)
|
||||||
|
|
||||||
// account for variable ip header length - incoming
|
// account for variable ip header length - incoming
|
||||||
h = ipv4.Header{
|
h = ipv4.Header{
|
||||||
@@ -116,7 +115,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
err = newPacket(buffer.Bytes(), true, p)
|
err = newPacket(buffer.Bytes(), true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A v6 packet with a hop-by-hop extension
|
// A v6 packet with a hop-by-hop extension
|
||||||
// ICMPv6 Payload (Echo Request)
|
// ICMPv6 Payload (Echo Request)
|
||||||
@@ -150,12 +149,12 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
// A full IPv6 header and 1 byte in the first extension, but missing
|
// A full IPv6 header and 1 byte in the first extension, but missing
|
||||||
// the length byte.
|
// the length byte.
|
||||||
err = newPacket(buffer.Bytes()[:41], true, p)
|
err = newPacket(buffer.Bytes()[:41], true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
// A full IPv6 header plus 1 full extension, but only 1 byte of the
|
||||||
// next layer, missing length byte
|
// next layer, missing length byte
|
||||||
err = newPacket(buffer.Bytes()[:49], true, p)
|
err = newPacket(buffer.Bytes()[:49], true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
err = nil
|
err = nil
|
||||||
|
|
||||||
// A good ICMP packet
|
// A good ICMP packet
|
||||||
@@ -218,7 +217,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
b[6] = 255 // 255 is a reserved protocol number
|
b[6] = 255 // 255 is a reserved protocol number
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
|
|
||||||
// A good UDP packet
|
// A good UDP packet
|
||||||
ip = layers.IPv6{
|
ip = layers.IPv6{
|
||||||
@@ -265,7 +264,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
|
|
||||||
// Too short UDP packet
|
// Too short UDP packet
|
||||||
err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes
|
err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// A good TCP packet
|
// A good TCP packet
|
||||||
b[6] = byte(layers.IPProtocolTCP)
|
b[6] = byte(layers.IPProtocolTCP)
|
||||||
@@ -292,7 +291,7 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
|
|
||||||
// Too short TCP packet
|
// Too short TCP packet
|
||||||
err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes
|
err = newPacket(b[:len(b)-10], false, p) // pull off the last 10 bytes
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// A good UDP packet with an AH header
|
// A good UDP packet with an AH header
|
||||||
ip = layers.IPv6{
|
ip = layers.IPv6{
|
||||||
@@ -337,12 +336,12 @@ func Test_newPacket_v6(t *testing.T) {
|
|||||||
|
|
||||||
// Ensure buffer bounds checking during processing
|
// Ensure buffer bounds checking during processing
|
||||||
err = newPacket(b[:41], true, p)
|
err = newPacket(b[:41], true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
|
|
||||||
// Invalid AH header
|
// Invalid AH header
|
||||||
b = buffer.Bytes()
|
b = buffer.Bytes()
|
||||||
err = newPacket(b, true, p)
|
err = newPacket(b, true, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6CouldNotFindPayload)
|
require.ErrorIs(t, err, ErrIPv6CouldNotFindPayload)
|
||||||
}
|
}
|
||||||
|
|
||||||
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
func Test_newPacket_ipv6Fragment(t *testing.T) {
|
||||||
@@ -449,7 +448,7 @@ func Test_newPacket_ipv6Fragment(t *testing.T) {
|
|||||||
|
|
||||||
// Too short of a fragment packet
|
// Too short of a fragment packet
|
||||||
err = newPacket(secondFrag[:len(secondFrag)-10], false, p)
|
err = newPacket(secondFrag[:len(secondFrag)-10], false, p)
|
||||||
require.ErrorIs(t, err, batch.ErrIPv6PacketTooShort)
|
require.ErrorIs(t, err, ErrIPv6PacketTooShort)
|
||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkParseV6(b *testing.B) {
|
func BenchmarkParseV6(b *testing.B) {
|
||||||
@@ -530,7 +529,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("Normal", func(b *testing.B) {
|
b.Run("Normal", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
if err = newPacket(normalPacket, true, fp); err != nil {
|
if err = parseV6(normalPacket, true, fp); err != nil {
|
||||||
b.Fatal(err)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -538,7 +537,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("FirstFragment", func(b *testing.B) {
|
b.Run("FirstFragment", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
if err = newPacket(firstFrag, true, fp); err != nil {
|
if err = parseV6(firstFrag, true, fp); err != nil {
|
||||||
b.Fatal(err)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -546,7 +545,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("SecondFragment", func(b *testing.B) {
|
b.Run("SecondFragment", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
if err = newPacket(secondFrag, true, fp); err != nil {
|
if err = parseV6(secondFrag, true, fp); err != nil {
|
||||||
b.Fatal(err)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -591,7 +590,7 @@ func BenchmarkParseV6(b *testing.B) {
|
|||||||
|
|
||||||
b.Run("200 HopByHop headers", func(b *testing.B) {
|
b.Run("200 HopByHop headers", func(b *testing.B) {
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
if err = newPacket(evilBytes, false, fp); err != nil {
|
if err = parseV6(evilBytes, false, fp); err != nil {
|
||||||
b.Fatal(err)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+17
-19
@@ -5,15 +5,8 @@ import "net/netip"
|
|||||||
type RxBatcher interface {
|
type RxBatcher interface {
|
||||||
// Reserve creates a pkt to borrow
|
// Reserve creates a pkt to borrow
|
||||||
Reserve(sz int) []byte
|
Reserve(sz int) []byte
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||||
// Walks IP+L4 headers itself; prefer CommitInbound when the caller already
|
|
||||||
// has an RxParsed in hand from ParsePacket.
|
|
||||||
Commit(pkt []byte) error
|
Commit(pkt []byte) error
|
||||||
// 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
|
// Flush emits every queued packet in arrival order. Returns the
|
||||||
// first error observed; keeps draining so one bad packet doesn't hold up
|
// first error observed; keeps draining so one bad packet doesn't hold up
|
||||||
// the rest. After Flush returns, borrowed payload slices may be recycled.
|
// the rest. After Flush returns, borrowed payload slices may be recycled.
|
||||||
@@ -21,15 +14,20 @@ type RxBatcher interface {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type TxBatcher interface {
|
type TxBatcher interface {
|
||||||
// Reserve creates a pkt to borrow
|
// Next returns a zero-length slice with slotCap capacity over the next unused
|
||||||
Reserve(sz int) []byte
|
// slot's backing bytes. The caller writes into the returned slice and then
|
||||||
// Commit borrows pkt and records its destination plus the 2-bit
|
// calls Commit with the final length and destination. Next returns nil when
|
||||||
// IP-level ECN codepoint to set on the outer (carrier) header. The
|
// the batch is full.
|
||||||
// caller must keep pkt valid until the next Flush. Pass 0 (Not-ECT)
|
Next() []byte
|
||||||
// to leave the outer ECN field unset.
|
// Commit records the slot just returned by Next as a packet of length n
|
||||||
Commit(pkt []byte, dst netip.AddrPort, outerECN byte)
|
// destined for dst.
|
||||||
// Flush emits every queued packet via the underlying batch writer in
|
Commit(n int, dst netip.AddrPort)
|
||||||
// arrival order. Returns an errors.Join of one or more errors. After Flush returns,
|
// Reset clears committed slots; backing storage is retained for reuse.
|
||||||
// borrowed payload slices may be recycled.
|
Reset()
|
||||||
Flush() error
|
// Len returns the number of committed packets.
|
||||||
|
Len() int
|
||||||
|
// Cap returns the maximum number of slots in the batch.
|
||||||
|
Cap() int
|
||||||
|
// Get returns the buffers needed to send the batch
|
||||||
|
Get() ([][]byte, []netip.AddrPort)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,163 +0,0 @@
|
|||||||
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]
|
|
||||||
}
|
|
||||||
@@ -1,443 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
@@ -1,394 +0,0 @@
|
|||||||
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))
|
|
||||||
}
|
|
||||||
@@ -1,174 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,133 +0,0 @@
|
|||||||
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...)
|
|
||||||
}
|
|
||||||
@@ -1,94 +0,0 @@
|
|||||||
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))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -40,13 +40,6 @@ func (p *Passthrough) Commit(pkt []byte) error {
|
|||||||
return nil
|
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 {
|
func (p *Passthrough) Flush() error {
|
||||||
var firstErr error
|
var firstErr error
|
||||||
for _, s := range p.slots {
|
for _, s := range p.slots {
|
||||||
@@ -55,7 +48,9 @@ func (p *Passthrough) Flush() error {
|
|||||||
firstErr = err
|
firstErr = err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
clear(p.slots)
|
for i := range p.slots {
|
||||||
|
p.slots[i] = nil
|
||||||
|
}
|
||||||
p.slots = p.slots[:0]
|
p.slots = p.slots[:0]
|
||||||
p.backing = p.backing[:0]
|
p.backing = p.backing[:0]
|
||||||
return firstErr
|
return firstErr
|
||||||
|
|||||||
+116
-341
@@ -4,9 +4,6 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
)
|
)
|
||||||
@@ -30,6 +27,18 @@ const tcpCoalesceMaxSegs = 64
|
|||||||
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
// into. IPv6 (40) + TCP with full options (60) = 100 bytes.
|
||||||
const tcpCoalesceHdrCap = 100
|
const tcpCoalesceHdrCap = 100
|
||||||
|
|
||||||
|
// initialSlots is the starting capacity of the slot pool. One flow per
|
||||||
|
// packet is the worst case so this matches a typical UDP recvmmsg batch.
|
||||||
|
const initialSlots = 64
|
||||||
|
|
||||||
|
// flowKey identifies a TCP flow by {src, dst, sport, dport, family}.
|
||||||
|
// Comparable, so linear scans over the slot list stay tight.
|
||||||
|
type flowKey struct {
|
||||||
|
src, dst [16]byte
|
||||||
|
sport, dport uint16
|
||||||
|
isV6 bool
|
||||||
|
}
|
||||||
|
|
||||||
// coalesceSlot is one entry in the coalescer's ordered event queue. When
|
// 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
|
// passthrough is true the slot holds a single borrowed packet that must be
|
||||||
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
// emitted verbatim (non-TCP, non-admissible TCP, or oversize seed). When
|
||||||
@@ -74,14 +83,7 @@ type TCPCoalescer struct {
|
|||||||
// removed from this map when they close (PSH or short-last-segment),
|
// removed from this map when they close (PSH or short-last-segment),
|
||||||
// when a non-admissible packet for that flow arrives, or in Flush.
|
// when a non-admissible packet for that flow arrives, or in Flush.
|
||||||
openSlots map[flowKey]*coalesceSlot
|
openSlots map[flowKey]*coalesceSlot
|
||||||
// lastSlot caches the most recently touched open slot. Steady-state
|
pool []*coalesceSlot // free list for reuse
|
||||||
// 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
|
backing []byte
|
||||||
}
|
}
|
||||||
@@ -94,7 +96,7 @@ func NewTCPCoalescer(w io.Writer) *TCPCoalescer {
|
|||||||
pool: make([]*coalesceSlot, 0, initialSlots),
|
pool: make([]*coalesceSlot, 0, initialSlots),
|
||||||
backing: make([]byte, 0, initialSlots*65535),
|
backing: make([]byte, 0, initialSlots*65535),
|
||||||
}
|
}
|
||||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoTCP); ok {
|
if gw, ok := w.(tio.GSOWriter); ok && gw.GSOSupported() {
|
||||||
c.gsoW = gw
|
c.gsoW = gw
|
||||||
}
|
}
|
||||||
return c
|
return c
|
||||||
@@ -118,13 +120,51 @@ type parsedTCP struct {
|
|||||||
// and IPv6 (no extension headers).
|
// and IPv6 (no extension headers).
|
||||||
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
||||||
var p parsedTCP
|
var p parsedTCP
|
||||||
ip, ok := parseIPPrologue(pkt, ipProtoTCP)
|
if len(pkt) < 20 {
|
||||||
if !ok {
|
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] != ipProtoTCP {
|
||||||
|
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])
|
||||||
|
pkt = pkt[:totalLen]
|
||||||
|
case 6:
|
||||||
|
if len(pkt) < 40 {
|
||||||
|
return p, false
|
||||||
|
}
|
||||||
|
if pkt[6] != ipProtoTCP {
|
||||||
|
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])
|
||||||
|
pkt = pkt[:40+payloadLen]
|
||||||
|
default:
|
||||||
return p, false
|
return p, false
|
||||||
}
|
}
|
||||||
pkt = ip.pkt
|
|
||||||
p.fk = ip.fk
|
|
||||||
p.ipHdrLen = ip.ipHdrLen
|
|
||||||
|
|
||||||
if len(pkt) < p.ipHdrLen+20 {
|
if len(pkt) < p.ipHdrLen+20 {
|
||||||
return p, false
|
return p, false
|
||||||
@@ -146,32 +186,27 @@ func parseTCPBase(pkt []byte) (parsedTCP, bool) {
|
|||||||
return p, true
|
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
|
// coalesceable reports whether a parsed TCP segment is eligible for
|
||||||
// coalescing. Accepts ACK, ACK|PSH, ACK|ECE, ACK|PSH|ECE with a
|
// coalescing. Accepts only ACK or ACK|PSH with a non-empty payload.
|
||||||
// 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 {
|
func (p parsedTCP) coalesceable() bool {
|
||||||
if p.flags&tcpFlagAck == 0 {
|
const ack = 0x10
|
||||||
return false
|
const psh = 0x08
|
||||||
}
|
if p.flags&^(ack|psh) != 0 || p.flags&ack == 0 {
|
||||||
if p.flags&^(tcpFlagAck|tcpFlagPsh|tcpFlagEce) != 0 {
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
return p.payLen > 0
|
return p.payLen > 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
func (c *TCPCoalescer) Reserve(sz int) []byte {
|
||||||
return reserveFromBacking(&c.backing, sz)
|
if len(c.backing)+sz > cap(c.backing) {
|
||||||
|
// Grow: allocate a fresh backing. Already-committed slices still
|
||||||
|
// reference the old array and remain valid until Flush drops them.
|
||||||
|
newCap := max(cap(c.backing)*2, sz)
|
||||||
|
c.backing = make([]byte, 0, newCap)
|
||||||
|
}
|
||||||
|
start := len(c.backing)
|
||||||
|
c.backing = c.backing[:start+sz]
|
||||||
|
return c.backing[start : start+sz : start+sz] //return zero length, sz-cap slice
|
||||||
}
|
}
|
||||||
|
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush,
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush,
|
||||||
@@ -182,77 +217,42 @@ func (c *TCPCoalescer) Commit(pkt []byte) error {
|
|||||||
c.addPassthrough(pkt)
|
c.addPassthrough(pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
info, ok := parseTCPBase(pkt)
|
info, ok := parseTCPBase(pkt)
|
||||||
if !ok {
|
if !ok {
|
||||||
c.addPassthrough(pkt)
|
// Non-TCP or malformed — can't possibly collide with an open flow.
|
||||||
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)
|
c.addPassthrough(pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if !info.coalesceable() {
|
if !info.coalesceable() {
|
||||||
// TCP but not admissible (SYN/FIN/RST/URG/CWR or zero-payload).
|
// TCP but not admissible (SYN/FIN/RST/URG/CWR/ECE or zero-payload).
|
||||||
// Seal this flow's open slot so later in-flow packets don't extend
|
// Seal this flow's open slot so later in-flow packets don't extend
|
||||||
// it and accidentally reorder past this passthrough.
|
// 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)
|
delete(c.openSlots, info.fk)
|
||||||
c.addPassthrough(pkt)
|
c.addPassthrough(pkt)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Single-flow fast path: with only one open flow the cache hits every
|
if open := c.openSlots[info.fk]; open != nil {
|
||||||
// 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) {
|
if c.canAppend(open, pkt, info) {
|
||||||
c.appendPayload(open, pkt, info)
|
c.appendPayload(open, pkt, info)
|
||||||
if open.psh {
|
if open.psh {
|
||||||
delete(c.openSlots, info.fk)
|
delete(c.openSlots, info.fk)
|
||||||
c.lastSlot = nil
|
|
||||||
} else {
|
|
||||||
c.lastSlot = open
|
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
// Can't extend — seal it and fall through to seed a fresh slot.
|
||||||
delete(c.openSlots, info.fk)
|
delete(c.openSlots, info.fk)
|
||||||
if c.lastSlot == open {
|
|
||||||
c.lastSlot = nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
c.seed(pkt, info)
|
c.seed(pkt, info)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush emits every queued event in (per-flow) seq order. Coalesced slots
|
// Flush emits every queued event in arrival order. Coalesced slots go out
|
||||||
// go out via WriteGSO; passthrough slots go out via plainW.Write.
|
// via WriteGSO; passthrough slots go out via plainW.Write. Returns the
|
||||||
// reorderForFlush first sorts each flow's slots into TCP-seq order within
|
// first error observed; keeps draining so one bad packet doesn't hold up
|
||||||
// passthrough-bounded segments and merges contiguous adjacent slots, so
|
// the rest. After Flush returns, borrowed payload slices may be recycled.
|
||||||
// any wire-side reorder that crossed an rxOrder batch boundary doesn't
|
|
||||||
// get amplified into kernel-visible reorder by the slot machinery.
|
|
||||||
// Returns the first error observed; keeps draining so one bad packet
|
|
||||||
// doesn't hold up the rest. After Flush returns, borrowed payload slices
|
|
||||||
// may be recycled.
|
|
||||||
func (c *TCPCoalescer) Flush() error {
|
func (c *TCPCoalescer) Flush() error {
|
||||||
c.reorderForFlush()
|
|
||||||
var first error
|
var first error
|
||||||
for _, s := range c.slots {
|
for _, s := range c.slots {
|
||||||
var err error
|
var err error
|
||||||
@@ -266,10 +266,13 @@ func (c *TCPCoalescer) Flush() error {
|
|||||||
}
|
}
|
||||||
c.release(s)
|
c.release(s)
|
||||||
}
|
}
|
||||||
clear(c.slots)
|
for i := range c.slots {
|
||||||
|
c.slots[i] = nil
|
||||||
|
}
|
||||||
c.slots = c.slots[:0]
|
c.slots = c.slots[:0]
|
||||||
clear(c.openSlots)
|
for k := range c.openSlots {
|
||||||
c.lastSlot = nil
|
delete(c.openSlots, k)
|
||||||
|
}
|
||||||
|
|
||||||
c.backing = c.backing[:0]
|
c.backing = c.backing[:0]
|
||||||
return first
|
return first
|
||||||
@@ -300,17 +303,11 @@ func (c *TCPCoalescer) seed(pkt []byte, info parsedTCP) {
|
|||||||
s.numSeg = 1
|
s.numSeg = 1
|
||||||
s.totalPay = info.payLen
|
s.totalPay = info.payLen
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
s.psh = info.flags&tcpFlagPsh != 0
|
s.psh = info.flags&0x08 != 0
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
||||||
c.slots = append(c.slots, s)
|
c.slots = append(c.slots, s)
|
||||||
if !s.psh {
|
if !s.psh {
|
||||||
c.openSlots[info.fk] = s
|
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
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -336,12 +333,6 @@ func (c *TCPCoalescer) canAppend(s *coalesceSlot, pkt []byte, info parsedTCP) bo
|
|||||||
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
if s.hdrLen+s.totalPay+info.payLen > tcpCoalesceBufSize {
|
||||||
return false
|
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) {
|
if !headersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -353,15 +344,12 @@ func (c *TCPCoalescer) appendPayload(s *coalesceSlot, pkt []byte, info parsedTCP
|
|||||||
s.numSeg++
|
s.numSeg++
|
||||||
s.totalPay += info.payLen
|
s.totalPay += info.payLen
|
||||||
s.nextSeq = info.seq + uint32(info.payLen)
|
s.nextSeq = info.seq + uint32(info.payLen)
|
||||||
if info.flags&tcpFlagPsh != 0 {
|
if info.flags&0x08 != 0 {
|
||||||
// Propagate PSH into the seed header so kernel TSO sets it on the
|
// Propagate PSH into the seed header so kernel TSO sets it on the
|
||||||
// last segment. Without this the sender's push signal is dropped.
|
// last segment. Without this the sender's push signal is dropped.
|
||||||
s.hdrBuf[s.ipHdrLen+13] |= tcpFlagPsh
|
s.hdrBuf[s.ipHdrLen+13] |= 0x08
|
||||||
}
|
}
|
||||||
// Merge IP-level CE marks into the seed: headersMatch ignores ECN, so
|
if info.payLen < s.gsoSize || info.flags&0x08 != 0 {
|
||||||
// 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
|
s.psh = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -379,7 +367,9 @@ func (c *TCPCoalescer) take() *coalesceSlot {
|
|||||||
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
func (c *TCPCoalescer) release(s *coalesceSlot) {
|
||||||
s.passthrough = false
|
s.passthrough = false
|
||||||
s.rawPkt = nil
|
s.rawPkt = nil
|
||||||
clear(s.payIovs)
|
for i := range s.payIovs {
|
||||||
|
s.payIovs[i] = nil
|
||||||
|
}
|
||||||
s.payIovs = s.payIovs[:0]
|
s.payIovs = s.payIovs[:0]
|
||||||
s.numSeg = 0
|
s.numSeg = 0
|
||||||
s.totalPay = 0
|
s.totalPay = 0
|
||||||
@@ -412,19 +402,37 @@ func (c *TCPCoalescer) flushSlot(s *coalesceSlot) error {
|
|||||||
tcsum := s.ipHdrLen + 16
|
tcsum := s.ipHdrLen + 16
|
||||||
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
binary.BigEndian.PutUint16(hdr[tcsum:tcsum+2], foldOnceNoInvert(psum))
|
||||||
|
|
||||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoTCP)
|
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs)
|
||||||
}
|
}
|
||||||
|
|
||||||
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
// headersMatch compares two IP+TCP header prefixes for byte-for-byte
|
||||||
// equality on every field that must be identical across coalesced
|
// equality on every field that must be identical across coalesced
|
||||||
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out, as is the
|
// segments. Size/IPID/IPCsum/seq/flags/tcpCsum are masked out.
|
||||||
// 2-bit IP-level ECN field — appendPayload merges CE into the seed.
|
|
||||||
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
||||||
if len(a) != len(b) {
|
if len(a) != len(b) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if !ipHeadersMatch(a, b, isV6) {
|
if isV6 {
|
||||||
return false
|
// IPv6: bytes [0:4] = version/TC/flow-label, [6:8] = next_hdr/hop,
|
||||||
|
// [8:40] = src+dst. Skip [4:6] payload length.
|
||||||
|
if !bytes.Equal(a[0:4], b[0:4]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:40], b[6:40]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// IPv4: [0:2] version/IHL/TOS, [6:10] flags/fragoff/TTL/proto,
|
||||||
|
// [12:20] src+dst. Skip [2:4] total len, [4:6] id, [10:12] csum.
|
||||||
|
if !bytes.Equal(a[0:2], b[0:2]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[6:10], b[6:10]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !bytes.Equal(a[12:20], b[12:20]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
}
|
}
|
||||||
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
// TCP: compare [0:4] ports, [8:13] ack+dataoff, [14:16] window,
|
||||||
// [18:tcpHdrLen] options (incl. urgent).
|
// [18:tcpHdrLen] options (incl. urgent).
|
||||||
@@ -444,238 +452,6 @@ func headersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|||||||
return true
|
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
|
// ipv4HdrChecksum computes the IPv4 header checksum over hdr (which must
|
||||||
// already have its checksum field zeroed) and returns the folded/inverted
|
// already have its checksum field zeroed) and returns the folded/inverted
|
||||||
// 16-bit value to store.
|
// 16-bit value to store.
|
||||||
@@ -693,10 +469,9 @@ func ipv4HdrChecksum(hdr []byte) uint16 {
|
|||||||
return ^uint16(sum)
|
return ^uint16(sum)
|
||||||
}
|
}
|
||||||
|
|
||||||
// pseudoSumIPv4 / pseudoSumIPv6 build the L4 pseudo-header partial sum
|
// pseudoSumIPv4 / pseudoSumIPv6 build the TCP pseudo-header partial sum
|
||||||
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
// expected by the virtio NEEDS_CSUM kernel path: the 32-bit accumulator
|
||||||
// before folding. proto selects the L4 (TCP or UDP); the UDP coalescer
|
// before folding.
|
||||||
// reuses these helpers.
|
|
||||||
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
func pseudoSumIPv4(src, dst []byte, proto byte, l4Len int) uint32 {
|
||||||
var sum uint32
|
var sum uint32
|
||||||
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
sum += uint32(binary.BigEndian.Uint16(src[0:2]))
|
||||||
|
|||||||
@@ -1,239 +0,0 @@
|
|||||||
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)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -3,8 +3,6 @@ package batch
|
|||||||
import (
|
import (
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// fakeTunWriter records plain Writes and WriteGSO calls without touching a
|
// fakeTunWriter records plain Writes and WriteGSO calls without touching a
|
||||||
@@ -52,7 +50,7 @@ func (w *fakeTunWriter) Write(p []byte) (int, error) {
|
|||||||
return len(p), nil
|
return len(p), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *fakeTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
func (w *fakeTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte) error {
|
||||||
hcopy := make([]byte, len(hdr)+len(transportHdr))
|
hcopy := make([]byte, len(hdr)+len(transportHdr))
|
||||||
copy(hcopy, hdr)
|
copy(hcopy, hdr)
|
||||||
copy(hcopy[len(hdr):], transportHdr)
|
copy(hcopy[len(hdr):], transportHdr)
|
||||||
@@ -77,9 +75,7 @@ func (w *fakeTunWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte,
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *fakeTunWriter) Capabilities() tio.Capabilities {
|
func (w *fakeTunWriter) GSOSupported() bool { return w.gsoEnabled }
|
||||||
return tio.Capabilities{TSO: w.gsoEnabled, USO: w.gsoEnabled}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTCPv4 constructs a minimal IPv4+TCP packet with the given payload,
|
// buildTCPv4 constructs a minimal IPv4+TCP packet with the given payload,
|
||||||
// seq, and flags. Assumes no IP options and a 20-byte TCP header.
|
// seq, and flags. Assumes no IP options and a 20-byte TCP header.
|
||||||
@@ -177,7 +173,7 @@ func TestCoalescerSeedThenFlushAlone(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// Single-segment flush goes through WriteGSO with GSO_NONE
|
// Single-segment flush now goes through WriteGSO with GSO_NONE
|
||||||
// (virtio NEEDS_CSUM lets the kernel fill in the L4 csum).
|
// (virtio NEEDS_CSUM lets the kernel fill in the L4 csum).
|
||||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
||||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
@@ -244,7 +240,7 @@ func TestCoalescerRejectsSeqGap(t *testing.T) {
|
|||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
// Each packet flushes as its own single-segment WriteGSO.
|
// Each packet flushes as its own single-segment WriteGSO now.
|
||||||
if len(w.gsoWrites) != 2 || len(w.writes) != 0 {
|
if len(w.gsoWrites) != 2 || len(w.writes) != 0 {
|
||||||
t.Fatalf("seq gap: want 2 gso writes got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
t.Fatalf("seq gap: want 2 gso writes got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
||||||
}
|
}
|
||||||
@@ -363,7 +359,7 @@ func TestCoalescerPropagatesPSHFromAppended(t *testing.T) {
|
|||||||
if err := c.Commit(buildTCPv4(2200, tcpAckPsh, pay)); err != nil {
|
if err := c.Commit(buildTCPv4(2200, tcpAckPsh, pay)); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := c.Flush(); err != nil {
|
if err := c.Flush(0); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(w.gsoWrites) != 1 {
|
if len(w.gsoWrites) != 1 {
|
||||||
@@ -548,14 +544,12 @@ func (w *orderedFakeWriter) Write(p []byte) (int, error) {
|
|||||||
return len(p), nil
|
return len(p), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *orderedFakeWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, _ tio.GSOProto) error {
|
func (w *orderedFakeWriter) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte) error {
|
||||||
w.events = append(w.events, "gso")
|
w.events = append(w.events, "gso")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *orderedFakeWriter) Capabilities() tio.Capabilities {
|
func (w *orderedFakeWriter) GSOSupported() bool { return w.gsoEnabled }
|
||||||
return tio.Capabilities{TSO: w.gsoEnabled, USO: w.gsoEnabled}
|
|
||||||
}
|
|
||||||
|
|
||||||
func stringSliceEq(a, b []string) bool {
|
func stringSliceEq(a, b []string) bool {
|
||||||
if len(a) != len(b) {
|
if len(a) != len(b) {
|
||||||
@@ -622,429 +616,3 @@ func TestCoalescerInterleavedFlowsPreserveOrdering(t *testing.T) {
|
|||||||
t.Errorf("unexpected segment counts: %v (want 2 and 3)", segCounts)
|
t.Errorf("unexpected segment counts: %v (want 2 and 3)", segCounts)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// ECN test helpers and constants.
|
|
||||||
|
|
||||||
const (
|
|
||||||
tcpEce = 0x40
|
|
||||||
tcpCwr = 0x80
|
|
||||||
|
|
||||||
// 2-bit IP-level ECN codepoints (lower 2 bits of IPv4 ToS / IPv6 TC).
|
|
||||||
ecnNotECT = 0x00
|
|
||||||
ecnECT1 = 0x01
|
|
||||||
ecnECT0 = 0x02
|
|
||||||
ecnCE = 0x03
|
|
||||||
)
|
|
||||||
|
|
||||||
// buildTCPv4WithToS is buildTCPv4 with caller-specified IPv4 ToS so tests can
|
|
||||||
// drive DSCP and ECN bits.
|
|
||||||
func buildTCPv4WithToS(tos byte, seq uint32, flags byte, payload []byte) []byte {
|
|
||||||
pkt := buildTCPv4(seq, flags, payload)
|
|
||||||
pkt[1] = tos
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTCPv6 mirrors buildTCPv4 for IPv6. tcLow is the low 4 bits of Traffic
|
|
||||||
// Class, which carries the ECN codepoint (mask 0x03) and the bottom 2 DSCP
|
|
||||||
// bits — enough to drive the ECN paths under test.
|
|
||||||
func buildTCPv6(tcLow byte, seq uint32, flags byte, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 40
|
|
||||||
const tcpHdrLen = 20
|
|
||||||
pkt := make([]byte, ipHdrLen+tcpHdrLen+len(payload))
|
|
||||||
|
|
||||||
pkt[0] = 0x60 // version=6, TC[7:4]=0
|
|
||||||
pkt[1] = (tcLow & 0x0f) << 4 // TC[3:0] in high nibble; flow=0
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpHdrLen+len(payload)))
|
|
||||||
pkt[6] = ipProtoTCP
|
|
||||||
pkt[7] = 64
|
|
||||||
copy(pkt[8:24], []byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1})
|
|
||||||
copy(pkt[24:40], []byte{0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2})
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], 1000)
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 2000)
|
|
||||||
binary.BigEndian.PutUint32(pkt[44:48], seq)
|
|
||||||
binary.BigEndian.PutUint32(pkt[48:52], 12345)
|
|
||||||
pkt[52] = 0x50
|
|
||||||
pkt[53] = flags
|
|
||||||
binary.BigEndian.PutUint16(pkt[54:56], 0xffff)
|
|
||||||
|
|
||||||
copy(pkt[60:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerCoalescesEceFlow confirms that ECN-Echo-marked ACKs (an
|
|
||||||
// ECN-aware flow under congestion) keep getting coalesced into a TSO
|
|
||||||
// superpacket instead of falling out to passthrough, and that the seed
|
|
||||||
// retains ECE on the wire.
|
|
||||||
func TestCoalescerCoalescesEceFlow(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
flags := byte(tcpAck | tcpEce)
|
|
||||||
if err := c.Commit(buildTCPv4(1000, flags, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(2200, flags, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if len(g.pays) != 2 {
|
|
||||||
t.Errorf("pay count=%d want 2", len(g.pays))
|
|
||||||
}
|
|
||||||
if seedFlags := g.hdr[20+13]; seedFlags&tcpEce == 0 {
|
|
||||||
t.Errorf("seed flags=0x%02x want ECE preserved", seedFlags)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerCwrSealsFlow confirms that a CWR-bearing segment in the
|
|
||||||
// middle of a flow goes to passthrough and seals the open slot, so a later
|
|
||||||
// in-flow segment seeds a new slot rather than extending the prior burst.
|
|
||||||
func TestCoalescerCwrSealsFlow(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(2200, tcpAck|tcpCwr, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want 1 plain write (CWR), got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
// Two GSO writes: the first seed before CWR, and a fresh seed after.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 1 {
|
|
||||||
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerEceMismatchReseeds confirms that toggling ECE mid-flow does
|
|
||||||
// not silently merge — receivers expect ECE either set on every segment of
|
|
||||||
// a CE-echoing window or none.
|
|
||||||
func TestCoalescerEceMismatchReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Commit(buildTCPv4(1000, tcpAck|tcpEce, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 separate seeds, got %d gso writes", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 1 {
|
|
||||||
t.Errorf("gso %d pay count=%d want 1", i, len(g.pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerMergesCEMark confirms that an ECT(0) burst with a single
|
|
||||||
// CE-marked packet still coalesces, and the merged superpacket carries CE.
|
|
||||||
func TestCoalescerMergesCEMark(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
if err := c.Commit(buildTCPv4WithToS(ecnECT0, 1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Router along the path stamped CE on this one.
|
|
||||||
if err := c.Commit(buildTCPv4WithToS(ecnCE, 2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4WithToS(ecnECT0, 3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 merged gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Errorf("pay count=%d want 3", len(g.pays))
|
|
||||||
}
|
|
||||||
if got := g.hdr[1] & 0x03; got != ecnCE {
|
|
||||||
t.Errorf("seed ECN=0x%02x want CE 0x%02x", got, ecnCE)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerDscpMismatchReseeds confirms that the new ECN-mask in
|
|
||||||
// headersMatch did not also relax DSCP — different DSCP must still split.
|
|
||||||
func TestCoalescerDscpMismatchReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
// Same ECN (Not-ECT), different DSCP (0x10 vs 0x20 in upper 6 bits).
|
|
||||||
tosA := byte(0x10<<2) | ecnNotECT
|
|
||||||
tosB := byte(0x20<<2) | ecnNotECT
|
|
||||||
if err := c.Commit(buildTCPv4WithToS(tosA, 1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4WithToS(tosB, 2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerIPv6CoalescesEceFlow is the IPv6 analogue of
|
|
||||||
// TestCoalescerCoalescesEceFlow.
|
|
||||||
func TestCoalescerIPv6CoalescesEceFlow(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
flags := byte(tcpAck | tcpEce)
|
|
||||||
if err := c.Commit(buildTCPv6(0, 1000, flags, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv6(0, 2200, flags, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if seedFlags := g.hdr[40+13]; seedFlags&tcpEce == 0 {
|
|
||||||
t.Errorf("seed flags=0x%02x want ECE preserved", seedFlags)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerSortsReorderedSeedsAndMerges feeds three same-flow MSS
|
|
||||||
// segments out of TCP-seq order (mimicking a wire reorder that escaped
|
|
||||||
// the rxOrder per-batch sort). Without the reorderForFlush sort+merge,
|
|
||||||
// each out-of-seq arrival would seed its own slot and the slots would
|
|
||||||
// emit in arrival order, producing a kernel-visible TCP reorder. With
|
|
||||||
// the sort+merge, the three slots are sorted by seq and folded back into
|
|
||||||
// one in-order TSO superpacket — same shape the receiver TCP would have
|
|
||||||
// seen had the wire never reordered.
|
|
||||||
func TestCoalescerSortsReorderedSeedsAndMerges(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
// Arrival order: seq 1000, 3400, 2200. The 3400 seeds a separate slot
|
|
||||||
// because 3400 != nextSeq=2200, then 2200 fails to extend the 3400 slot
|
|
||||||
// and seeds its own. Three slots end up in c.slots; reorderForFlush
|
|
||||||
// should sort them into [1000,2200,3400] and merge them back into one.
|
|
||||||
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 merged gso write got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Fatalf("merged segs=%d want 3", len(g.pays))
|
|
||||||
}
|
|
||||||
const ipHdrLen = 20
|
|
||||||
if seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8]); seedSeq != 1000 {
|
|
||||||
t.Errorf("merged seed seq=%d want 1000 (lowest)", seedSeq)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerSortAcrossFlowsMergesEachIndependently checks that two
|
|
||||||
// flows interleaved with reorder are each sorted-and-merged in isolation
|
|
||||||
// without any cross-flow contamination.
|
|
||||||
func TestCoalescerSortAcrossFlowsMergesEachIndependently(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
// Flow A (sport 1000) seq 100, 1300; flow B (sport 3000) seq 500, 1700.
|
|
||||||
// Arrival: A.1300, B.1700, A.100, B.500 — every flow reordered.
|
|
||||||
if err := c.Commit(buildTCPv4Ports(1000, 2000, 1300, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4Ports(3000, 2000, 1700, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4Ports(1000, 2000, 100, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4Ports(3000, 2000, 500, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (one per flow merged), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 2 {
|
|
||||||
t.Errorf("gso[%d] segs=%d want 2", i, len(g.pays))
|
|
||||||
}
|
|
||||||
const ipHdrLen = 20
|
|
||||||
seedSeq := binary.BigEndian.Uint32(g.hdr[ipHdrLen+4 : ipHdrLen+8])
|
|
||||||
sport := binary.BigEndian.Uint16(g.hdr[ipHdrLen : ipHdrLen+2])
|
|
||||||
// Each flow's merged seed should be the LOWER of its two seqs.
|
|
||||||
switch sport {
|
|
||||||
case 1000:
|
|
||||||
if seedSeq != 100 {
|
|
||||||
t.Errorf("flow A seed seq=%d want 100", seedSeq)
|
|
||||||
}
|
|
||||||
case 3000:
|
|
||||||
if seedSeq != 500 {
|
|
||||||
t.Errorf("flow B seed seq=%d want 500", seedSeq)
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
t.Errorf("unexpected sport %d", sport)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerSortKeepsPSHBoundary verifies that a PSH-sealed slot is
|
|
||||||
// not folded into a later seq-contiguous slot — PSH placement is part of
|
|
||||||
// the wire signal and merging across it would shift the receiver's push
|
|
||||||
// boundary by an arbitrary number of segments.
|
|
||||||
func TestCoalescerSortKeepsPSHBoundary(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
// Seq 1000 (no PSH) + 2200 (PSH) → seal one slot with PSH set.
|
|
||||||
// Seq 3400 (no PSH) is contiguous to 3400 from seq 2200+1200; without
|
|
||||||
// the PSH check it would merge in.
|
|
||||||
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(2200, tcpAckPsh, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (PSH-sealed and fresh seed), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerSortKeepsPassthroughBarrier confirms a passthrough slot in
|
|
||||||
// the middle of the queue prevents the post-sort merge from folding
|
|
||||||
// across it. Reordered same-flow data on either side of the passthrough
|
|
||||||
// is sorted/merged independently.
|
|
||||||
func TestCoalescerSortKeepsPassthroughBarrier(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
// First two segments seed S1 (then a 3400 reorder seeds S2).
|
|
||||||
if err := c.Commit(buildTCPv4(1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv4(3400, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Non-coalesceable packet (SYN+ACK) flushes S1's openSlots entry and
|
|
||||||
// becomes a passthrough barrier in c.slots.
|
|
||||||
if err := c.Commit(buildTCPv4(9999, tcpSyn|tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Post-barrier same-flow data: should never end up before the SYN.
|
|
||||||
if err := c.Commit(buildTCPv4(2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// We expect: gso(merged 1000+3400 ranges sorted but not contiguous so 2
|
|
||||||
// gso writes), plain(SYN), gso(2200 alone). The pre-barrier sort should
|
|
||||||
// land 1000 before 3400, and the post-barrier 2200 stays after the SYN.
|
|
||||||
if len(w.writes) != 1 {
|
|
||||||
t.Fatalf("want 1 plain SYN passthrough, got %d", len(w.writes))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestCoalescerIPv6MergesCEMark is the IPv6 analogue of
|
|
||||||
// TestCoalescerMergesCEMark. ECN bits live in TC[1:0] = byte 1 mask 0x30.
|
|
||||||
func TestCoalescerIPv6MergesCEMark(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewTCPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
// tcLow is the low 4 bits of TC; ECN occupies the bottom 2 of those.
|
|
||||||
if err := c.Commit(buildTCPv6(ecnECT0, 1000, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildTCPv6(ecnCE, 2200, tcpAck, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 merged gso write, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
// Byte 1 high nibble holds TC[3:0]; ECN is the low 2 bits of that nibble,
|
|
||||||
// which appears in byte 1 mask 0x30 (>>4 to read the codepoint value).
|
|
||||||
if got := (g.hdr[1] >> 4) & 0x03; got != ecnCE {
|
|
||||||
t.Errorf("seed v6 ECN=0x%02x want CE 0x%02x", got, ecnCE)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSortRunZeroAllocs(t *testing.T) {
|
|
||||||
c := &TCPCoalescer{}
|
|
||||||
mk := func(srcByte byte, seq uint32, pay int) *coalesceSlot {
|
|
||||||
s := &coalesceSlot{nextSeq: seq + uint32(pay), totalPay: pay}
|
|
||||||
s.fk.src[0] = srcByte
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
run := []*coalesceSlot{
|
|
||||||
mk(3, 5000, 100),
|
|
||||||
mk(1, 1000, 50),
|
|
||||||
mk(2, 2000, 75),
|
|
||||||
mk(1, 900, 50),
|
|
||||||
mk(3, 4900, 100),
|
|
||||||
mk(2, 1925, 75),
|
|
||||||
mk(1, 1050, 50),
|
|
||||||
mk(3, 5100, 100),
|
|
||||||
}
|
|
||||||
|
|
||||||
allocs := testing.AllocsPerRun(100, func() {
|
|
||||||
// Re-shuffle so each run actually does sorting work.
|
|
||||||
run[0], run[1], run[2], run[3] = run[3], run[2], run[1], run[0]
|
|
||||||
c.sortRun(run)
|
|
||||||
})
|
|
||||||
if allocs != 0 {
|
|
||||||
t.Fatalf("sortRun allocates %v times per run; want 0", allocs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+38
-41
@@ -4,61 +4,58 @@ import "net/netip"
|
|||||||
|
|
||||||
const SendBatchCap = 128
|
const SendBatchCap = 128
|
||||||
|
|
||||||
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
// SendBatch accumulates encrypted UDP packets for potential TX offloading.
|
||||||
type batchWriter interface {
|
|
||||||
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
|
||||||
}
|
|
||||||
|
|
||||||
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
|
||||||
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
||||||
// The backing arena grows on demand: when there isn't room for the next slot
|
// The backing storage holds up to batchCap packets of slotCap bytes each;
|
||||||
// we allocate a fresh backing array. Already-committed slices keep referencing
|
// bufs and dsts are parallel slices of committed slots.
|
||||||
// the old array and remain valid until Flush drops them.
|
|
||||||
type SendBatch struct {
|
type SendBatch struct {
|
||||||
out batchWriter
|
bufs [][]byte
|
||||||
bufs [][]byte
|
dsts []netip.AddrPort
|
||||||
dsts []netip.AddrPort
|
backing []byte
|
||||||
ecns []byte
|
slotCap int
|
||||||
backing []byte
|
batchCap int
|
||||||
|
nextSlot int
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewSendBatch(out batchWriter, batchCap, slotCap int) *SendBatch {
|
func NewSendBatch(batchCap, slotCap int) *SendBatch {
|
||||||
return &SendBatch{
|
return &SendBatch{
|
||||||
out: out,
|
bufs: make([][]byte, 0, batchCap),
|
||||||
bufs: make([][]byte, 0, batchCap),
|
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||||
dsts: make([]netip.AddrPort, 0, batchCap),
|
backing: make([]byte, batchCap*slotCap),
|
||||||
ecns: make([]byte, 0, batchCap),
|
slotCap: slotCap,
|
||||||
backing: make([]byte, 0, batchCap*slotCap),
|
batchCap: batchCap,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *SendBatch) Reserve(sz int) []byte {
|
func (b *SendBatch) Next() []byte {
|
||||||
if len(b.backing)+sz > cap(b.backing) {
|
if b.nextSlot >= b.batchCap {
|
||||||
// Grow: allocate a fresh backing. Already-committed slices still
|
return nil
|
||||||
// reference the old array and remain valid until Flush drops them.
|
|
||||||
newCap := max(cap(b.backing)*2, sz)
|
|
||||||
b.backing = make([]byte, 0, newCap)
|
|
||||||
}
|
}
|
||||||
start := len(b.backing)
|
start := b.nextSlot * b.slotCap
|
||||||
b.backing = b.backing[:start+sz]
|
return b.backing[start : start : start+b.slotCap] //set len to 0 but cap to slotCap
|
||||||
return b.backing[start : start+sz : start+sz]
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
|
func (b *SendBatch) Commit(n int, dst netip.AddrPort) {
|
||||||
b.bufs = append(b.bufs, pkt)
|
start := b.nextSlot * b.slotCap
|
||||||
|
b.bufs = append(b.bufs, b.backing[start:start+n])
|
||||||
b.dsts = append(b.dsts, dst)
|
b.dsts = append(b.dsts, dst)
|
||||||
b.ecns = append(b.ecns, outerECN)
|
b.nextSlot++
|
||||||
}
|
}
|
||||||
|
|
||||||
func (b *SendBatch) Flush() error {
|
func (b *SendBatch) Reset() {
|
||||||
var err error
|
|
||||||
if len(b.bufs) > 0 {
|
|
||||||
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
|
||||||
}
|
|
||||||
clear(b.bufs)
|
|
||||||
b.bufs = b.bufs[:0]
|
b.bufs = b.bufs[:0]
|
||||||
b.dsts = b.dsts[:0]
|
b.dsts = b.dsts[:0]
|
||||||
b.ecns = b.ecns[:0]
|
b.nextSlot = 0
|
||||||
b.backing = b.backing[:0]
|
}
|
||||||
return err
|
|
||||||
|
func (b *SendBatch) Len() int {
|
||||||
|
return len(b.bufs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Cap() int {
|
||||||
|
return b.batchCap
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Get() ([][]byte, []netip.AddrPort) {
|
||||||
|
return b.bufs, b.dsts
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,120 +5,65 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
type fakeBatchWriter struct {
|
func TestSendBatchBookkeeping(t *testing.T) {
|
||||||
bufs [][]byte
|
b := NewSendBatch(4, 32)
|
||||||
addrs []netip.AddrPort
|
if b.Len() != 0 || b.Cap() != 4 {
|
||||||
ecns []byte
|
t.Fatalf("fresh batch: len=%d cap=%d", b.Len(), b.Cap())
|
||||||
}
|
|
||||||
|
|
||||||
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
|
||||||
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
|
||||||
// returns, so tests must capture data before that happens.
|
|
||||||
w.bufs = make([][]byte, len(bufs))
|
|
||||||
for i, b := range bufs {
|
|
||||||
cp := make([]byte, len(b))
|
|
||||||
copy(cp, b)
|
|
||||||
w.bufs[i] = cp
|
|
||||||
}
|
}
|
||||||
w.addrs = append(w.addrs[:0], addrs...)
|
|
||||||
w.ecns = append(w.ecns[:0], ecns...)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
|
||||||
fw := &fakeBatchWriter{}
|
|
||||||
b := NewSendBatch(fw, 4, 32)
|
|
||||||
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||||
for i := 0; i < 4; i++ {
|
for i := 0; i < 4; i++ {
|
||||||
slot := b.Reserve(32)
|
slot := b.Next()
|
||||||
if cap(slot) != 32 {
|
if slot == nil {
|
||||||
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
t.Fatalf("slot %d: Next returned nil before cap", i)
|
||||||
}
|
}
|
||||||
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
if cap(slot) != 32 || len(slot) != 0 {
|
||||||
b.Commit(pkt, ap, 0)
|
t.Fatalf("slot %d: got len=%d cap=%d want len=0 cap=32", i, len(slot), cap(slot))
|
||||||
|
}
|
||||||
|
// Write a marker byte.
|
||||||
|
slot = append(slot, byte(i), byte(i+1), byte(i+2))
|
||||||
|
b.Commit(len(slot), ap)
|
||||||
}
|
}
|
||||||
if err := b.Flush(); err != nil {
|
if b.Next() != nil {
|
||||||
t.Fatalf("Flush: %v", err)
|
t.Fatalf("Next should return nil when full")
|
||||||
}
|
}
|
||||||
if len(fw.bufs) != 4 {
|
if b.Len() != 4 {
|
||||||
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
t.Fatalf("Len=%d want 4", b.Len())
|
||||||
}
|
}
|
||||||
for i, buf := range fw.bufs {
|
for i, buf := range b.bufs {
|
||||||
if len(buf) != 3 || buf[0] != byte(i) {
|
if len(buf) != 3 || buf[0] != byte(i) {
|
||||||
t.Errorf("buf %d: %x", i, buf)
|
t.Errorf("buf %d: %x", i, buf)
|
||||||
}
|
}
|
||||||
if fw.addrs[i] != ap {
|
if b.dsts[i] != ap {
|
||||||
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
t.Errorf("dst %d: got %v want %v", i, b.dsts[i], ap)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush again with nothing committed — should be a no-op.
|
// Reset returns empty and Next works again.
|
||||||
fw.bufs = nil
|
b.Reset()
|
||||||
if err := b.Flush(); err != nil {
|
if b.Len() != 0 {
|
||||||
t.Fatalf("empty Flush: %v", err)
|
t.Fatalf("after Reset Len=%d want 0", b.Len())
|
||||||
}
|
}
|
||||||
if fw.bufs != nil {
|
slot := b.Next()
|
||||||
t.Fatalf("empty Flush triggered WriteBatch")
|
if slot == nil || cap(slot) != 32 {
|
||||||
}
|
t.Fatalf("after Reset Next nil or wrong cap: %v cap=%d", slot == nil, cap(slot))
|
||||||
|
|
||||||
// Reuse after Flush.
|
|
||||||
slot := b.Reserve(32)
|
|
||||||
if cap(slot) != 32 {
|
|
||||||
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||||
fw := &fakeBatchWriter{}
|
b := NewSendBatch(3, 8)
|
||||||
b := NewSendBatch(fw, 3, 8)
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||||
|
|
||||||
|
// Fill three slots, each with its own sentinel byte.
|
||||||
for i := 0; i < 3; i++ {
|
for i := 0; i < 3; i++ {
|
||||||
s := b.Reserve(8)
|
s := b.Next()
|
||||||
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
s = append(s, byte(0xA0+i), byte(0xB0+i))
|
||||||
b.Commit(pkt, ap, 0)
|
b.Commit(len(s), ap)
|
||||||
}
|
|
||||||
if err := b.Flush(); err != nil {
|
|
||||||
t.Fatalf("Flush: %v", err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for i, buf := range fw.bufs {
|
for i, buf := range b.bufs {
|
||||||
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
||||||
t.Errorf("slot %d corrupted: %x", i, buf)
|
t.Errorf("slot %d corrupted: %x", i, buf)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
|
||||||
fw := &fakeBatchWriter{}
|
|
||||||
// Tiny initial backing forces a grow on the second Reserve.
|
|
||||||
b := NewSendBatch(fw, 1, 4)
|
|
||||||
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
|
||||||
|
|
||||||
s1 := b.Reserve(4)
|
|
||||||
pkt1 := append(s1[:0], 0x11, 0x22, 0x33, 0x44)
|
|
||||||
b.Commit(pkt1, ap, 0)
|
|
||||||
|
|
||||||
s2 := b.Reserve(8) // exceeds remaining cap, triggers grow
|
|
||||||
pkt2 := append(s2[:0], 0xA, 0xB, 0xC, 0xD, 0xE)
|
|
||||||
b.Commit(pkt2, ap, 0)
|
|
||||||
|
|
||||||
// pkt1 must still be intact even though backing reallocated.
|
|
||||||
if pkt1[0] != 0x11 || pkt1[3] != 0x44 {
|
|
||||||
t.Fatalf("first packet corrupted by grow: %x", pkt1)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := b.Flush(); err != nil {
|
|
||||||
t.Fatalf("Flush: %v", err)
|
|
||||||
}
|
|
||||||
if len(fw.bufs) != 2 {
|
|
||||||
t.Fatalf("got %d bufs want 2", len(fw.bufs))
|
|
||||||
}
|
|
||||||
if fw.bufs[0][0] != 0x11 || fw.bufs[0][3] != 0x44 {
|
|
||||||
t.Errorf("first packet on the wire: %x", fw.bufs[0])
|
|
||||||
}
|
|
||||||
if fw.bufs[1][0] != 0xA || fw.bufs[1][4] != 0xE {
|
|
||||||
t.Errorf("second packet on the wire: %x", fw.bufs[1])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,336 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio"
|
|
||||||
)
|
|
||||||
|
|
||||||
// ipProtoUDP is the IANA protocol number for UDP.
|
|
||||||
const ipProtoUDP = 17
|
|
||||||
|
|
||||||
// udpCoalesceBufSize caps total bytes per UDP superpacket. Mirrors the
|
|
||||||
// kernel's gso_max_size; payloads beyond this are emitted as-is.
|
|
||||||
const udpCoalesceBufSize = 65535
|
|
||||||
|
|
||||||
// udpCoalesceMaxSegs caps how many segments we'll coalesce. Kernel UDP-GSO
|
|
||||||
// accepts up to 64 segments per skb (UDP_MAX_SEGMENTS); stay under that.
|
|
||||||
const udpCoalesceMaxSegs = 64
|
|
||||||
|
|
||||||
// udpCoalesceHdrCap is the scratch space we copy a seed's IP+UDP header
|
|
||||||
// into. IPv6 (40) + UDP (8) = 48; round up for safety.
|
|
||||||
const udpCoalesceHdrCap = 64
|
|
||||||
|
|
||||||
// udpSlot is one entry in the UDPCoalescer's ordered event queue. Same
|
|
||||||
// passthrough-vs-coalesced shape as the TCP coalescer's slot, but no
|
|
||||||
// seq/PSH/CWR bookkeeping — UDP segments only need 5-tuple + length
|
|
||||||
// matching to coalesce.
|
|
||||||
type udpSlot struct {
|
|
||||||
passthrough bool
|
|
||||||
rawPkt []byte // borrowed when passthrough
|
|
||||||
|
|
||||||
fk flowKey
|
|
||||||
hdrBuf [udpCoalesceHdrCap]byte
|
|
||||||
hdrLen int
|
|
||||||
ipHdrLen int
|
|
||||||
isV6 bool
|
|
||||||
gsoSize int // per-segment UDP payload length
|
|
||||||
numSeg int
|
|
||||||
totalPay int
|
|
||||||
// sealed closes the chain: set when a sub-gsoSize segment is appended
|
|
||||||
// (kernel UDP-GSO requires every segment but the last to be exactly
|
|
||||||
// gsoSize) or when limits are hit. No further appends after.
|
|
||||||
sealed bool
|
|
||||||
payIovs [][]byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// UDPCoalescer accumulates adjacent in-flow UDP datagrams across multiple
|
|
||||||
// concurrent flows and emits each flow's run as a single GSO_UDP_L4
|
|
||||||
// superpacket via tio.GSOWriter. Falls back to per-packet writes when the
|
|
||||||
// underlying writer doesn't support USO.
|
|
||||||
//
|
|
||||||
// All output — coalesced or not — is deferred until Flush so per-flow
|
|
||||||
// arrival order is preserved on the wire. Cross-flow order is NOT preserved
|
|
||||||
// across the TCP/UDP/passthrough split when this coalescer runs alongside
|
|
||||||
// others — see multi_coalesce.go. Per-flow order is preserved because a
|
|
||||||
// single 5-tuple only ever lands in one lane and each lane preserves its
|
|
||||||
// own slot order.
|
|
||||||
//
|
|
||||||
// Owns no locks; one coalescer per TUN write queue.
|
|
||||||
type UDPCoalescer struct {
|
|
||||||
plainW io.Writer
|
|
||||||
gsoW tio.GSOWriter // nil when the queue can't accept GSO_UDP_L4
|
|
||||||
|
|
||||||
slots []*udpSlot
|
|
||||||
openSlots map[flowKey]*udpSlot
|
|
||||||
pool []*udpSlot
|
|
||||||
|
|
||||||
backing []byte
|
|
||||||
}
|
|
||||||
|
|
||||||
// NewUDPCoalescer wraps w. The caller is responsible for only constructing
|
|
||||||
// this when the underlying Queue's Capabilities advertise USO; otherwise
|
|
||||||
// the kernel may reject GSO_UDP_L4 writes. If w does not implement
|
|
||||||
// tio.GSOWriter at all (single-packet Queue), the coalescer degrades to
|
|
||||||
// plain Writes — same defensive shape as the TCP coalescer.
|
|
||||||
func NewUDPCoalescer(w io.Writer) *UDPCoalescer {
|
|
||||||
c := &UDPCoalescer{
|
|
||||||
plainW: w,
|
|
||||||
slots: make([]*udpSlot, 0, initialSlots),
|
|
||||||
openSlots: make(map[flowKey]*udpSlot, initialSlots),
|
|
||||||
pool: make([]*udpSlot, 0, initialSlots),
|
|
||||||
backing: make([]byte, 0, initialSlots*udpCoalesceBufSize),
|
|
||||||
}
|
|
||||||
if gw, ok := tio.SupportsGSO(w, tio.GSOProtoUDP); ok {
|
|
||||||
c.gsoW = gw
|
|
||||||
}
|
|
||||||
return c
|
|
||||||
}
|
|
||||||
|
|
||||||
// parsedUDP holds the fields extracted from a single parse so later steps
|
|
||||||
// (admission, slot lookup, canAppend) don't re-walk the header.
|
|
||||||
type parsedUDP struct {
|
|
||||||
fk flowKey
|
|
||||||
ipHdrLen int
|
|
||||||
hdrLen int // ipHdrLen + 8
|
|
||||||
payLen int
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseUDP extracts the flow key and IP/UDP offsets for a UDP packet.
|
|
||||||
// Returns ok=false for non-UDP, malformed, or unsupported header shapes
|
|
||||||
// (IPv4 with options/fragmentation, IPv6 with extension headers).
|
|
||||||
func parseUDP(pkt []byte) (parsedUDP, bool) {
|
|
||||||
var p parsedUDP
|
|
||||||
ip, ok := parseIPPrologue(pkt, ipProtoUDP)
|
|
||||||
if !ok {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
pkt = ip.pkt
|
|
||||||
p.fk = ip.fk
|
|
||||||
p.ipHdrLen = ip.ipHdrLen
|
|
||||||
|
|
||||||
if len(pkt) < p.ipHdrLen+8 {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.hdrLen = p.ipHdrLen + 8
|
|
||||||
// UDP `length` field: must equal IP-derived length-of-UDP-header-plus-payload.
|
|
||||||
udpLen := int(binary.BigEndian.Uint16(pkt[p.ipHdrLen+4 : p.ipHdrLen+6]))
|
|
||||||
if udpLen < 8 || udpLen > len(pkt)-p.ipHdrLen {
|
|
||||||
return p, false
|
|
||||||
}
|
|
||||||
p.payLen = udpLen - 8
|
|
||||||
p.fk.sport = binary.BigEndian.Uint16(pkt[p.ipHdrLen : p.ipHdrLen+2])
|
|
||||||
p.fk.dport = binary.BigEndian.Uint16(pkt[p.ipHdrLen+2 : p.ipHdrLen+4])
|
|
||||||
return p, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) Reserve(sz int) []byte {
|
|
||||||
return reserveFromBacking(&c.backing, sz)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Commit borrows pkt. The caller must keep pkt valid until the next Flush.
|
|
||||||
func (c *UDPCoalescer) Commit(pkt []byte) error {
|
|
||||||
if c.gsoW == nil {
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
info, ok := parseUDP(pkt)
|
|
||||||
if !ok {
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return c.commitParsed(pkt, info)
|
|
||||||
}
|
|
||||||
|
|
||||||
// commitParsed is the post-parse half of Commit. The caller must have
|
|
||||||
// already verified parseUDP succeeded. Used by MultiCoalescer.Commit to
|
|
||||||
// avoid re-walking the IP/UDP header.
|
|
||||||
func (c *UDPCoalescer) commitParsed(pkt []byte, info parsedUDP) error {
|
|
||||||
if c.gsoW == nil {
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if open := c.openSlots[info.fk]; open != nil {
|
|
||||||
if c.canAppend(open, pkt, info) {
|
|
||||||
c.appendPayload(open, pkt, info)
|
|
||||||
if open.sealed {
|
|
||||||
delete(c.openSlots, info.fk)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
// Can't extend — seal it and fall through to seed a fresh slot.
|
|
||||||
delete(c.openSlots, info.fk)
|
|
||||||
}
|
|
||||||
c.seed(pkt, info)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) Flush() error {
|
|
||||||
var first error
|
|
||||||
for _, s := range c.slots {
|
|
||||||
var err error
|
|
||||||
if s.passthrough {
|
|
||||||
_, err = c.plainW.Write(s.rawPkt)
|
|
||||||
} else {
|
|
||||||
err = c.flushSlot(s)
|
|
||||||
}
|
|
||||||
if err != nil && first == nil {
|
|
||||||
first = err
|
|
||||||
}
|
|
||||||
c.release(s)
|
|
||||||
}
|
|
||||||
clear(c.slots)
|
|
||||||
c.slots = c.slots[:0]
|
|
||||||
clear(c.openSlots)
|
|
||||||
c.backing = c.backing[:0]
|
|
||||||
return first
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) addPassthrough(pkt []byte) {
|
|
||||||
s := c.take()
|
|
||||||
s.passthrough = true
|
|
||||||
s.rawPkt = pkt
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) seed(pkt []byte, info parsedUDP) {
|
|
||||||
if info.hdrLen > udpCoalesceHdrCap || info.hdrLen+info.payLen > udpCoalesceBufSize {
|
|
||||||
c.addPassthrough(pkt)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
s := c.take()
|
|
||||||
s.passthrough = false
|
|
||||||
s.rawPkt = nil
|
|
||||||
copy(s.hdrBuf[:], pkt[:info.hdrLen])
|
|
||||||
s.hdrLen = info.hdrLen
|
|
||||||
s.ipHdrLen = info.ipHdrLen
|
|
||||||
s.isV6 = info.fk.isV6
|
|
||||||
s.fk = info.fk
|
|
||||||
s.gsoSize = info.payLen
|
|
||||||
s.numSeg = 1
|
|
||||||
s.totalPay = info.payLen
|
|
||||||
s.sealed = false
|
|
||||||
s.payIovs = append(s.payIovs[:0], pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
c.slots = append(c.slots, s)
|
|
||||||
c.openSlots[info.fk] = s
|
|
||||||
}
|
|
||||||
|
|
||||||
// canAppend reports whether info's packet extends the slot's seed.
|
|
||||||
// Kernel UDP-GSO requires every segment except possibly the last to be
|
|
||||||
// exactly gsoSize, and the last may be shorter (≤ gsoSize).
|
|
||||||
func (c *UDPCoalescer) canAppend(s *udpSlot, pkt []byte, info parsedUDP) bool {
|
|
||||||
if s.sealed {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.hdrLen != s.hdrLen {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.numSeg >= udpCoalesceMaxSegs {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if info.payLen > s.gsoSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if s.hdrLen+s.totalPay+info.payLen > udpCoalesceBufSize {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !udpHeadersMatch(s.hdrBuf[:s.hdrLen], pkt[:info.hdrLen], s.isV6, s.ipHdrLen) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) appendPayload(s *udpSlot, pkt []byte, info parsedUDP) {
|
|
||||||
s.payIovs = append(s.payIovs, pkt[info.hdrLen:info.hdrLen+info.payLen])
|
|
||||||
s.numSeg++
|
|
||||||
s.totalPay += info.payLen
|
|
||||||
// Merge IP-level CE marks into the seed (same trick TCP coalescer uses).
|
|
||||||
mergeECNIntoSeed(s.hdrBuf[:s.ipHdrLen], pkt[:s.ipHdrLen], s.isV6)
|
|
||||||
if info.payLen < s.gsoSize {
|
|
||||||
// Last-segment-can-be-shorter: this seals the chain.
|
|
||||||
s.sealed = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) take() *udpSlot {
|
|
||||||
if n := len(c.pool); n > 0 {
|
|
||||||
s := c.pool[n-1]
|
|
||||||
c.pool[n-1] = nil
|
|
||||||
c.pool = c.pool[:n-1]
|
|
||||||
return s
|
|
||||||
}
|
|
||||||
return &udpSlot{}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *UDPCoalescer) release(s *udpSlot) {
|
|
||||||
s.passthrough = false
|
|
||||||
s.rawPkt = nil
|
|
||||||
clear(s.payIovs)
|
|
||||||
s.payIovs = s.payIovs[:0]
|
|
||||||
s.numSeg = 0
|
|
||||||
s.totalPay = 0
|
|
||||||
s.sealed = false
|
|
||||||
c.pool = append(c.pool, s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// flushSlot patches the IP header total length / IPv6 payload length and
|
|
||||||
// the UDP length to the *total* across all coalesced segments, then seeds
|
|
||||||
// the UDP checksum field with the pseudo-header partial (single-fold, not
|
|
||||||
// inverted) per virtio NEEDS_CSUM. The kernel's ip_rcv_core (v4) and
|
|
||||||
// ip6_rcv_core (v6) trim the skb to those length fields, so per-segment
|
|
||||||
// values would silently drop everything but the first segment. The kernel
|
|
||||||
// then walks each segment in __udp_gso_segment, recomputing per-segment
|
|
||||||
// uh->len / iph->tot_len / IPv6 plen and adjusting the checksum via
|
|
||||||
// `check = csum16_add(csum16_sub(uh->check, uh->len), newlen)` — meaning
|
|
||||||
// our seed's uh->check must be consistent with the seed's uh->len, which
|
|
||||||
// is what passing the total to both pseudoSum and the UDP length field
|
|
||||||
// guarantees.
|
|
||||||
func (c *UDPCoalescer) flushSlot(s *udpSlot) error {
|
|
||||||
hdr := s.hdrBuf[:s.hdrLen]
|
|
||||||
total := s.hdrLen + s.totalPay // full IP+UDP+all_payloads bytes
|
|
||||||
l4Len := total - s.ipHdrLen // total UDP (8 + sum of payloads)
|
|
||||||
|
|
||||||
if s.isV6 {
|
|
||||||
binary.BigEndian.PutUint16(hdr[4:6], uint16(l4Len))
|
|
||||||
} else {
|
|
||||||
binary.BigEndian.PutUint16(hdr[2:4], uint16(total))
|
|
||||||
hdr[10] = 0
|
|
||||||
hdr[11] = 0
|
|
||||||
binary.BigEndian.PutUint16(hdr[10:12], ipv4HdrChecksum(hdr[:s.ipHdrLen]))
|
|
||||||
}
|
|
||||||
|
|
||||||
// UDP length field (offset 4 inside the UDP header) = total UDP size.
|
|
||||||
binary.BigEndian.PutUint16(hdr[s.ipHdrLen+4:s.ipHdrLen+6], uint16(l4Len))
|
|
||||||
|
|
||||||
var psum uint32
|
|
||||||
if s.isV6 {
|
|
||||||
psum = pseudoSumIPv6(hdr[8:24], hdr[24:40], ipProtoUDP, l4Len)
|
|
||||||
} else {
|
|
||||||
psum = pseudoSumIPv4(hdr[12:16], hdr[16:20], ipProtoUDP, l4Len)
|
|
||||||
}
|
|
||||||
udpCsumOff := s.ipHdrLen + 6
|
|
||||||
binary.BigEndian.PutUint16(hdr[udpCsumOff:udpCsumOff+2], foldOnceNoInvert(psum))
|
|
||||||
|
|
||||||
return c.gsoW.WriteGSO(hdr[:s.ipHdrLen], hdr[s.ipHdrLen:], s.payIovs, tio.GSOProtoUDP)
|
|
||||||
}
|
|
||||||
|
|
||||||
// udpHeadersMatch compares two IP+UDP header prefixes for byte-equality on
|
|
||||||
// every field that must be identical across coalesced segments. Length
|
|
||||||
// fields and the ECN bits in IP TOS/TC are masked out — appendPayload
|
|
||||||
// merges CE into the seed; flushSlot rewrites lengths.
|
|
||||||
func udpHeadersMatch(a, b []byte, isV6 bool, ipHdrLen int) bool {
|
|
||||||
if len(a) != len(b) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if !ipHeadersMatch(a, b, isV6) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// UDP: compare sport+dport ([0:4]). Skip length [4:6] and checksum [6:8] —
|
|
||||||
// length varies (we rewrite at flush) and the checksum will be redone.
|
|
||||||
udp := ipHdrLen
|
|
||||||
if a[udp] != b[udp] || a[udp+1] != b[udp+1] || a[udp+2] != b[udp+2] || a[udp+3] != b[udp+3] {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
@@ -1,383 +0,0 @@
|
|||||||
package batch
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
// buildUDPv4 builds a minimal IPv4+UDP packet with the given payload and ports.
|
|
||||||
func buildUDPv4(sport, dport uint16, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 20
|
|
||||||
const udpHdrLen = 8
|
|
||||||
total := ipHdrLen + udpHdrLen + len(payload)
|
|
||||||
pkt := make([]byte, total)
|
|
||||||
|
|
||||||
pkt[0] = 0x45
|
|
||||||
pkt[1] = 0x00
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(total))
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], 0)
|
|
||||||
binary.BigEndian.PutUint16(pkt[6:8], 0x4000)
|
|
||||||
pkt[8] = 64
|
|
||||||
pkt[9] = ipProtoUDP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], sport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], dport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[24:26], uint16(udpHdrLen+len(payload)))
|
|
||||||
binary.BigEndian.PutUint16(pkt[26:28], 0)
|
|
||||||
|
|
||||||
copy(pkt[28:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildUDPv6 builds a minimal IPv6+UDP packet.
|
|
||||||
func buildUDPv6(sport, dport uint16, payload []byte) []byte {
|
|
||||||
const ipHdrLen = 40
|
|
||||||
const udpHdrLen = 8
|
|
||||||
total := ipHdrLen + udpHdrLen + len(payload)
|
|
||||||
pkt := make([]byte, total)
|
|
||||||
|
|
||||||
pkt[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpHdrLen+len(payload)))
|
|
||||||
pkt[6] = ipProtoUDP
|
|
||||||
pkt[7] = 64
|
|
||||||
pkt[8] = 0xfe
|
|
||||||
pkt[9] = 0x80
|
|
||||||
pkt[23] = 1
|
|
||||||
pkt[24] = 0xfe
|
|
||||||
pkt[25] = 0x80
|
|
||||||
pkt[39] = 2
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], sport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], dport)
|
|
||||||
binary.BigEndian.PutUint16(pkt[44:46], uint16(udpHdrLen+len(payload)))
|
|
||||||
binary.BigEndian.PutUint16(pkt[46:48], 0)
|
|
||||||
|
|
||||||
copy(pkt[48:], payload)
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerPassthroughWhenGSOUnavailable(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: false}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 100))
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 0 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("no Add-time writes: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("want single plain write, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerNonUDPPassthrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
// ICMP packet
|
|
||||||
pkt := make([]byte, 28)
|
|
||||||
pkt[0] = 0x45
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], 28)
|
|
||||||
pkt[9] = 1
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("ICMP must pass through unchanged: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerSeedThenFlushAlone(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 800))
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Single-segment flush goes through WriteGSO; the writer infers GSO_NONE
|
|
||||||
// from len(pays)==1 and the kernel fills in the UDP csum (NEEDS_CSUM).
|
|
||||||
if len(w.gsoWrites) != 1 || len(w.writes) != 0 {
|
|
||||||
t.Fatalf("single-seg flush: writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUDPCoalescerCoalescesEqualSized(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
for i := 0; i < 3; i++ {
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if g.gsoSize != 1200 {
|
|
||||||
t.Errorf("gsoSize=%d want 1200", g.gsoSize)
|
|
||||||
}
|
|
||||||
if len(g.pays) != 3 {
|
|
||||||
t.Errorf("pay count=%d want 3", len(g.pays))
|
|
||||||
}
|
|
||||||
if g.csumStart != 20 {
|
|
||||||
t.Errorf("csumStart=%d want 20", g.csumStart)
|
|
||||||
}
|
|
||||||
// IP totalLen and UDP length must be the TOTAL across all segments —
|
|
||||||
// the kernel's ip_rcv_core trims skbs to iph->tot_len, so a per-segment
|
|
||||||
// value would silently drop everything but the first segment. Total =
|
|
||||||
// IP(20) + UDP(8) + 3*1200 = 3628.
|
|
||||||
gotTotalLen := binary.BigEndian.Uint16(g.hdr[2:4])
|
|
||||||
if gotTotalLen != 3628 {
|
|
||||||
t.Errorf("ipv4 total_len=%d want 3628 (must be total across segments)", gotTotalLen)
|
|
||||||
}
|
|
||||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[20+4 : 20+6])
|
|
||||||
if gotUDPLen != 8+3*1200 {
|
|
||||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Last segment may be shorter, sealing the chain.
|
|
||||||
func TestUDPCoalescerShortLastSegmentSeals(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
full := make([]byte, 1200)
|
|
||||||
tail := make([]byte, 600)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, tail)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// A 4th packet, even same-sized, must NOT join — chain is sealed.
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, full)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (sealed + new seed), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[0].pays) != 3 {
|
|
||||||
t.Errorf("first super: want 3 pays, got %d", len(w.gsoWrites[0].pays))
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[1].pays) != 1 {
|
|
||||||
t.Errorf("second super: want 1 pay (re-seed), got %d", len(w.gsoWrites[1].pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// A larger-than-gsoSize packet cannot extend the slot — it reseeds.
|
|
||||||
func TestUDPCoalescerLargerThanSeedReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 800))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, make([]byte, 1200))); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 separate seeds, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Different 5-tuples must not coalesce.
|
|
||||||
func TestUDPCoalescerDifferentFlowsKeepSeparate(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pay := make([]byte, 800)
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(buildUDPv4(2000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// Two flows × 2 datagrams each = 2 superpackets of 2 segments.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (one per flow), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
for i, g := range w.gsoWrites {
|
|
||||||
if len(g.pays) != 2 {
|
|
||||||
t.Errorf("super %d: want 2 pays, got %d", i, len(g.pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Caps at udpCoalesceMaxSegs.
|
|
||||||
func TestUDPCoalescerCapsAtMaxSegs(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pay := make([]byte, 100)
|
|
||||||
for i := 0; i < udpCoalesceMaxSegs+5; i++ {
|
|
||||||
if err := c.Commit(buildUDPv4(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
// First superpacket holds udpCoalesceMaxSegs segments; the spillover
|
|
||||||
// reseeds a new one.
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 gso writes (cap then reseed), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[0].pays) != udpCoalesceMaxSegs {
|
|
||||||
t.Errorf("first super: pays=%d want %d", len(w.gsoWrites[0].pays), udpCoalesceMaxSegs)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites[1].pays) != 5 {
|
|
||||||
t.Errorf("second super: pays=%d want 5", len(w.gsoWrites[1].pays))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// CE marks on appended segments must be merged into the seed's IP TOS.
|
|
||||||
func TestUDPCoalescerMergesCEMark(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pay := make([]byte, 800)
|
|
||||||
pkt0 := buildUDPv4(1000, 53, pay) // ECN=00
|
|
||||||
pkt1 := buildUDPv4(1000, 53, pay)
|
|
||||||
pkt1[1] = 0x03 // CE
|
|
||||||
pkt2 := buildUDPv4(1000, 53, pay)
|
|
||||||
if err := c.Commit(pkt0); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(pkt1); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(pkt2); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 merged gso write, got %d (plain=%d)", len(w.gsoWrites), len(w.writes))
|
|
||||||
}
|
|
||||||
if w.gsoWrites[0].hdr[1]&0x03 != 0x03 {
|
|
||||||
t.Errorf("CE not merged into seed (tos=%#x)", w.gsoWrites[0].hdr[1])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// IPv6 path: same flow, equal-sized → coalesced.
|
|
||||||
func TestUDPCoalescerIPv6Coalesces(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pay := make([]byte, 1200)
|
|
||||||
for i := 0; i < 3; i++ {
|
|
||||||
if err := c.Commit(buildUDPv6(1000, 53, pay)); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 1 {
|
|
||||||
t.Fatalf("want 1 gso write, got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
g := w.gsoWrites[0]
|
|
||||||
if !g.isV6 {
|
|
||||||
t.Errorf("expected v6 write")
|
|
||||||
}
|
|
||||||
if g.csumStart != 40 {
|
|
||||||
t.Errorf("csumStart=%d want 40", g.csumStart)
|
|
||||||
}
|
|
||||||
// IPv6 payload_len and UDP length must be TOTAL — kernel's
|
|
||||||
// ip6_rcv_core trims to payload_len + ipv6 hdr size. Total UDP = 8 +
|
|
||||||
// 3*1200 = 3608.
|
|
||||||
gotPlen := binary.BigEndian.Uint16(g.hdr[4:6])
|
|
||||||
if gotPlen != 8+3*1200 {
|
|
||||||
t.Errorf("ipv6 payload_len=%d want %d (must be total)", gotPlen, 8+3*1200)
|
|
||||||
}
|
|
||||||
gotUDPLen := binary.BigEndian.Uint16(g.hdr[40+4 : 40+6])
|
|
||||||
if gotUDPLen != 8+3*1200 {
|
|
||||||
t.Errorf("udp len=%d want %d", gotUDPLen, 8+3*1200)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// DSCP differences must reseed (headers don't match outside ECN).
|
|
||||||
func TestUDPCoalescerDSCPMismatchReseeds(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pay := make([]byte, 800)
|
|
||||||
pkt0 := buildUDPv4(1000, 53, pay)
|
|
||||||
pkt1 := buildUDPv4(1000, 53, pay)
|
|
||||||
pkt1[1] = 0xb8 // EF DSCP, ECN=0
|
|
||||||
if err := c.Commit(pkt0); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Commit(pkt1); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.gsoWrites) != 2 {
|
|
||||||
t.Fatalf("want 2 separate seeds (different DSCP), got %d", len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Fragmented IPv4 must not be coalesced.
|
|
||||||
func TestUDPCoalescerFragmentedIPv4PassesThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
|
||||||
binary.BigEndian.PutUint16(pkt[6:8], 0x2000) // MF=1
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("frag must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// IPv4 with options is not admissible (we require IHL=5).
|
|
||||||
func TestUDPCoalescerIPv4WithOptionsPassesThrough(t *testing.T) {
|
|
||||||
w := &fakeTunWriter{gsoEnabled: true}
|
|
||||||
c := NewUDPCoalescer(w)
|
|
||||||
pkt := buildUDPv4(1000, 53, make([]byte, 200))
|
|
||||||
pkt[0] = 0x46 // IHL = 6 (24-byte IPv4 header — has options)
|
|
||||||
if err := c.Commit(pkt); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := c.Flush(); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if len(w.writes) != 1 || len(w.gsoWrites) != 0 {
|
|
||||||
t.Fatalf("ipv4-with-options must pass through plain, got writes=%d gso=%d", len(w.writes), len(w.gsoWrites))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,23 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
import (
|
|
||||||
"golang.org/x/sys/cpu"
|
|
||||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
func checksumAVX2(buf []byte, initial uint16) uint16
|
|
||||||
|
|
||||||
var hasAVX2 = cpu.X86.HasAVX2
|
|
||||||
|
|
||||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
|
||||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum that
|
|
||||||
// dispatches to a hand-written AVX2 routine on amd64 CPUs that support it,
|
|
||||||
// falling back to gvisor's pure-Go implementation otherwise. The result
|
|
||||||
// matches gvisor's bit-for-bit for any buffer length and initial seed.
|
|
||||||
func Checksum(buf []byte, initial uint16) uint16 {
|
|
||||||
if hasAVX2 {
|
|
||||||
return checksumAVX2(buf, initial)
|
|
||||||
}
|
|
||||||
return gvisorchecksum.Checksum(buf, initial)
|
|
||||||
}
|
|
||||||
@@ -1,157 +0,0 @@
|
|||||||
#include "textflag.h"
|
|
||||||
|
|
||||||
// func checksumAVX2(buf []byte, initial uint16) uint16
|
|
||||||
//
|
|
||||||
// Computes the RFC 1071 ones-complement sum of buf, seeded with initial.
|
|
||||||
//
|
|
||||||
// Algorithm: sum the buffer treating it as a stream of uint32s in machine
|
|
||||||
// (little-endian) byte order, accumulating into 64-bit lanes (top 32 bits
|
|
||||||
// hold cross-add carries — at 1 byte / lane / iter we have 32 bits of
|
|
||||||
// headroom which is far more than the 16 KB/64 KB max practical inputs).
|
|
||||||
// At the end we fold to 16 bits and byte-swap once to recover the on-wire
|
|
||||||
// (big-endian) result. RFC 1071 §1.2.B byte-order independence makes this
|
|
||||||
// equivalent to summing as 16-bit big-endian words.
|
|
||||||
//
|
|
||||||
// The ymm accumulators (Y4..Y7) hold 4 uint64 lanes each = 16 parallel
|
|
||||||
// partial sums. The main loop loads 64 bytes per iter as four 16-byte
|
|
||||||
// chunks, zero-extending each chunk's four uint32s into a ymm via
|
|
||||||
// VPMOVZXDQ-from-memory, then VPADDQ into a separate accumulator per
|
|
||||||
// chunk to break the dep chain. After the vector loop the lane sums are
|
|
||||||
// horizontally reduced and merged with a scalar accumulator that handles
|
|
||||||
// the trailing 0..63 bytes plus the (byte-swapped) initial seed.
|
|
||||||
TEXT ·checksumAVX2(SB), NOSPLIT, $0-34
|
|
||||||
MOVQ buf_base+0(FP), SI
|
|
||||||
MOVQ buf_len+8(FP), CX
|
|
||||||
MOVWQZX initial+24(FP), AX
|
|
||||||
|
|
||||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
|
||||||
// with the rest of the accumulator. The final fold's bswap16 will undo
|
|
||||||
// this and convert the whole result back to BE.
|
|
||||||
XCHGB AH, AL
|
|
||||||
|
|
||||||
CMPQ CX, $32
|
|
||||||
JLT scalar_tail
|
|
||||||
|
|
||||||
VPXOR Y4, Y4, Y4
|
|
||||||
VPXOR Y5, Y5, Y5
|
|
||||||
VPXOR Y6, Y6, Y6
|
|
||||||
VPXOR Y7, Y7, Y7
|
|
||||||
|
|
||||||
CMPQ CX, $64
|
|
||||||
JLT loop32
|
|
||||||
|
|
||||||
loop64:
|
|
||||||
VPMOVZXDQ (SI), Y0
|
|
||||||
VPMOVZXDQ 16(SI), Y1
|
|
||||||
VPMOVZXDQ 32(SI), Y2
|
|
||||||
VPMOVZXDQ 48(SI), Y3
|
|
||||||
VPADDQ Y0, Y4, Y4
|
|
||||||
VPADDQ Y1, Y5, Y5
|
|
||||||
VPADDQ Y2, Y6, Y6
|
|
||||||
VPADDQ Y3, Y7, Y7
|
|
||||||
ADDQ $64, SI
|
|
||||||
SUBQ $64, CX
|
|
||||||
CMPQ CX, $64
|
|
||||||
JGE loop64
|
|
||||||
|
|
||||||
loop32:
|
|
||||||
CMPQ CX, $32
|
|
||||||
JLT reduce_vec
|
|
||||||
VPMOVZXDQ (SI), Y0
|
|
||||||
VPMOVZXDQ 16(SI), Y1
|
|
||||||
VPADDQ Y0, Y4, Y4
|
|
||||||
VPADDQ Y1, Y5, Y5
|
|
||||||
ADDQ $32, SI
|
|
||||||
SUBQ $32, CX
|
|
||||||
JMP loop32
|
|
||||||
|
|
||||||
reduce_vec:
|
|
||||||
// Combine the four ymm accumulators into Y4.
|
|
||||||
VPADDQ Y5, Y4, Y4
|
|
||||||
VPADDQ Y7, Y6, Y6
|
|
||||||
VPADDQ Y6, Y4, Y4
|
|
||||||
|
|
||||||
// Horizontally reduce Y4's four uint64 lanes to a single scalar.
|
|
||||||
VEXTRACTI128 $1, Y4, X5
|
|
||||||
VPADDQ X5, X4, X4
|
|
||||||
VPSHUFD $0x4e, X4, X5
|
|
||||||
VPADDQ X5, X4, X4
|
|
||||||
VMOVQ X4, R8
|
|
||||||
VZEROUPPER
|
|
||||||
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
|
|
||||||
scalar_tail:
|
|
||||||
// Handle remaining 0..63 bytes (or the entire buffer if it was < 32).
|
|
||||||
CMPQ CX, $8
|
|
||||||
JLT tail4
|
|
||||||
|
|
||||||
loop8:
|
|
||||||
ADDQ (SI), AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
ADDQ $8, SI
|
|
||||||
SUBQ $8, CX
|
|
||||||
CMPQ CX, $8
|
|
||||||
JGE loop8
|
|
||||||
|
|
||||||
tail4:
|
|
||||||
CMPQ CX, $4
|
|
||||||
JLT tail2
|
|
||||||
MOVL (SI), R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
ADDQ $4, SI
|
|
||||||
SUBQ $4, CX
|
|
||||||
|
|
||||||
tail2:
|
|
||||||
CMPQ CX, $2
|
|
||||||
JLT tail1
|
|
||||||
MOVWQZX (SI), R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
ADDQ $2, SI
|
|
||||||
SUBQ $2, CX
|
|
||||||
|
|
||||||
tail1:
|
|
||||||
TESTQ CX, CX
|
|
||||||
JZ fold
|
|
||||||
MOVBQZX (SI), R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
ADCQ $0, AX
|
|
||||||
|
|
||||||
fold:
|
|
||||||
// Fold the 64-bit accumulator to 16 bits via four rounds, mirroring
|
|
||||||
// gvisor's reduce(). Each pair (split, add) halves the live width;
|
|
||||||
// the truncation steps absorb the single bit that may be left over
|
|
||||||
// after each add so the next round's bound holds.
|
|
||||||
|
|
||||||
// 64 → 33 bits.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $32, R8
|
|
||||||
MOVL AX, AX
|
|
||||||
ADDQ R8, AX
|
|
||||||
|
|
||||||
// 33 → 32 bits. AX += (AX>>32); truncate to 32. AX is now ≤ 0xFFFF_FFFF.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $32, R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
MOVL AX, AX
|
|
||||||
|
|
||||||
// 32 → 17 bits.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $16, R8
|
|
||||||
MOVWQZX AX, AX
|
|
||||||
ADDQ R8, AX
|
|
||||||
|
|
||||||
// 17 → 16 bits. AX += (AX>>16); the trailing MOVW truncates bit 16.
|
|
||||||
MOVQ AX, R8
|
|
||||||
SHRQ $16, R8
|
|
||||||
ADDQ R8, AX
|
|
||||||
|
|
||||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
|
||||||
// to big-endian to match the gvisor API contract.
|
|
||||||
XCHGB AH, AL
|
|
||||||
|
|
||||||
MOVW AX, ret+32(FP)
|
|
||||||
RET
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
func checksumNEON(buf []byte, initial uint16) uint16
|
|
||||||
|
|
||||||
// Checksum computes the RFC 1071 ones-complement sum of buf, seeded with
|
|
||||||
// initial. It is a drop-in replacement for gvisor's checksum.Checksum
|
|
||||||
// that dispatches to a hand-written NEON routine. NEON is mandatory in
|
|
||||||
// armv8 so no feature check is needed.
|
|
||||||
func Checksum(buf []byte, initial uint16) uint16 {
|
|
||||||
return checksumNEON(buf, initial)
|
|
||||||
}
|
|
||||||
@@ -1,143 +0,0 @@
|
|||||||
#include "textflag.h"
|
|
||||||
|
|
||||||
// func checksumNEON(buf []byte, initial uint16) uint16
|
|
||||||
//
|
|
||||||
// Mirrors the algorithm in checksum_amd64.s: sum the buffer treating it as
|
|
||||||
// a stream of uint32s in machine (little-endian) byte order, accumulating
|
|
||||||
// into 64-bit lanes that have ample carry headroom; fold and byte-swap once
|
|
||||||
// at the very end to recover the on-wire (big-endian) result.
|
|
||||||
//
|
|
||||||
// Each loop iteration loads 64 bytes via VLD1.P into V0..V3 (4 Q regs).
|
|
||||||
// VUADDW takes the low two uint32 lanes of a Q reg, zero-extends them to
|
|
||||||
// uint64, and adds them into a 2×uint64 accumulator; VUADDW2 does the same
|
|
||||||
// for the high two lanes. Four ymm-equivalent accumulators (V8..V11) get
|
|
||||||
// updated twice per iter to break the dep chain. Tail bytes go through a
|
|
||||||
// scalar ADCS chain seeded with the byte-swapped initial.
|
|
||||||
TEXT ·checksumNEON(SB), NOSPLIT, $0-34
|
|
||||||
MOVD buf_base+0(FP), R0
|
|
||||||
MOVD buf_len+8(FP), R1
|
|
||||||
MOVHU initial+24(FP), R2
|
|
||||||
|
|
||||||
// Pre-byteswap initial into the LE-summing space so it merges directly
|
|
||||||
// with the rest of the accumulator.
|
|
||||||
REV16W R2, R2
|
|
||||||
|
|
||||||
MOVD ZR, R3 // scalar accumulator
|
|
||||||
|
|
||||||
CMP $32, R1
|
|
||||||
BLT scalar_tail
|
|
||||||
|
|
||||||
VEOR V8.B16, V8.B16, V8.B16
|
|
||||||
VEOR V9.B16, V9.B16, V9.B16
|
|
||||||
VEOR V10.B16, V10.B16, V10.B16
|
|
||||||
VEOR V11.B16, V11.B16, V11.B16
|
|
||||||
|
|
||||||
CMP $64, R1
|
|
||||||
BLT loop16_init
|
|
||||||
|
|
||||||
loop64:
|
|
||||||
VLD1.P 64(R0), [V0.B16, V1.B16, V2.B16, V3.B16]
|
|
||||||
VUADDW V0.S2, V8.D2, V8.D2
|
|
||||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
|
||||||
VUADDW V1.S2, V10.D2, V10.D2
|
|
||||||
VUADDW2 V1.S4, V11.D2, V11.D2
|
|
||||||
VUADDW V2.S2, V8.D2, V8.D2
|
|
||||||
VUADDW2 V2.S4, V9.D2, V9.D2
|
|
||||||
VUADDW V3.S2, V10.D2, V10.D2
|
|
||||||
VUADDW2 V3.S4, V11.D2, V11.D2
|
|
||||||
SUB $64, R1, R1
|
|
||||||
CMP $64, R1
|
|
||||||
BGE loop64
|
|
||||||
|
|
||||||
loop16_init:
|
|
||||||
CMP $16, R1
|
|
||||||
BLT reduce_vec
|
|
||||||
|
|
||||||
loop16:
|
|
||||||
VLD1.P 16(R0), [V0.B16]
|
|
||||||
VUADDW V0.S2, V8.D2, V8.D2
|
|
||||||
VUADDW2 V0.S4, V9.D2, V9.D2
|
|
||||||
SUB $16, R1, R1
|
|
||||||
CMP $16, R1
|
|
||||||
BGE loop16
|
|
||||||
|
|
||||||
reduce_vec:
|
|
||||||
// Combine the four accumulators into V8.
|
|
||||||
VADD V9.D2, V8.D2, V8.D2
|
|
||||||
VADD V11.D2, V10.D2, V10.D2
|
|
||||||
VADD V10.D2, V8.D2, V8.D2
|
|
||||||
|
|
||||||
// Horizontal-add the two lanes of V8.D2 into a single uint64.
|
|
||||||
VADDP V8.D2, V8.D2, V8.D2
|
|
||||||
VMOV V8.D[0], R8
|
|
||||||
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
|
|
||||||
scalar_tail:
|
|
||||||
CMP $8, R1
|
|
||||||
BLT tail4
|
|
||||||
|
|
||||||
loop8:
|
|
||||||
MOVD.P 8(R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
SUB $8, R1, R1
|
|
||||||
CMP $8, R1
|
|
||||||
BGE loop8
|
|
||||||
|
|
||||||
tail4:
|
|
||||||
CMP $4, R1
|
|
||||||
BLT tail2
|
|
||||||
MOVWU.P 4(R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
SUB $4, R1, R1
|
|
||||||
|
|
||||||
tail2:
|
|
||||||
CMP $2, R1
|
|
||||||
BLT tail1
|
|
||||||
MOVHU.P 2(R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
SUB $2, R1, R1
|
|
||||||
|
|
||||||
tail1:
|
|
||||||
CBZ R1, fold
|
|
||||||
MOVBU (R0), R8
|
|
||||||
ADDS R8, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
|
|
||||||
fold:
|
|
||||||
// Merge the byte-swapped initial into our LE-form accumulator.
|
|
||||||
ADDS R2, R3, R3
|
|
||||||
ADC ZR, R3, R3
|
|
||||||
|
|
||||||
// 64 → 33 bits.
|
|
||||||
LSR $32, R3, R8
|
|
||||||
AND $0xffffffff, R3, R3
|
|
||||||
ADD R8, R3, R3
|
|
||||||
|
|
||||||
// 33 → 32 (truncate after adding bit 32 back).
|
|
||||||
LSR $32, R3, R8
|
|
||||||
ADD R8, R3, R3
|
|
||||||
AND $0xffffffff, R3, R3
|
|
||||||
|
|
||||||
// 32 → 17.
|
|
||||||
LSR $16, R3, R8
|
|
||||||
AND $0xffff, R3, R3
|
|
||||||
ADD R8, R3, R3
|
|
||||||
|
|
||||||
// 17 → 16 (truncation absorbs bit 16 below).
|
|
||||||
LSR $16, R3, R8
|
|
||||||
ADD R8, R3, R3
|
|
||||||
|
|
||||||
// AX low 16 bits hold the 16-bit sum in machine (LE) byte order; flip
|
|
||||||
// to big-endian to match the gvisor API contract. REV16W swaps bytes
|
|
||||||
// within each 16-bit halfword of the low 32 bits, so it acts as a
|
|
||||||
// 16-bit byte-swap on the live low 16.
|
|
||||||
REV16W R3, R3
|
|
||||||
AND $0xffff, R3, R3
|
|
||||||
|
|
||||||
MOVH R3, ret+32(FP)
|
|
||||||
RET
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
//go:build !amd64 && !arm64
|
|
||||||
|
|
||||||
package checksum
|
|
||||||
|
|
||||||
import gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
|
|
||||||
// Checksum delegates to gvisor on architectures without a hand-written body.
|
|
||||||
func Checksum(buf []byte, initial uint16) uint16 {
|
|
||||||
return gvisorchecksum.Checksum(buf, initial)
|
|
||||||
}
|
|
||||||
@@ -1,190 +0,0 @@
|
|||||||
package checksum
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"math/rand/v2"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
gvisorchecksum "gvisor.dev/gvisor/pkg/tcpip/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
// TestChecksumMatchesGvisor walks lengths from 0 to 4096, with several initial
|
|
||||||
// seeds and a handful of starting alignments, asserting that our local
|
|
||||||
// Checksum matches gvisor's reference bit-for-bit.
|
|
||||||
func TestChecksumMatchesGvisor(t *testing.T) {
|
|
||||||
rng := rand.New(rand.NewPCG(1, 2))
|
|
||||||
const padFront = 16
|
|
||||||
|
|
||||||
// Random pool large enough for the longest case + alignment slop.
|
|
||||||
pool := make([]byte, 4096+padFront)
|
|
||||||
for i := range pool {
|
|
||||||
pool[i] = byte(rng.Uint32())
|
|
||||||
}
|
|
||||||
|
|
||||||
seeds := []uint16{0, 0x0001, 0xabcd, 0xffff, 0x1234, 0xfedc}
|
|
||||||
offsets := []int{0, 1, 2, 3, 4, 5, 7, 8, 15, 16}
|
|
||||||
|
|
||||||
for length := 0; length <= 4096; length++ {
|
|
||||||
for _, seed := range seeds {
|
|
||||||
for _, off := range offsets {
|
|
||||||
if off+length > len(pool) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
buf := pool[off : off+length]
|
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
|
||||||
got := Checksum(buf, seed)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("len=%d off=%d seed=%#x: got %#04x want %#04x",
|
|
||||||
length, off, seed, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestChecksumPatternedBuffers exercises specific byte patterns that have
|
|
||||||
// historically tripped up checksum implementations: all-zero, all-0xff,
|
|
||||||
// alternating, and ascending sequences.
|
|
||||||
func TestChecksumPatternedBuffers(t *testing.T) {
|
|
||||||
for length := 0; length <= 256; length++ {
|
|
||||||
patterns := map[string][]byte{
|
|
||||||
"zeros": make([]byte, length),
|
|
||||||
"ones": bytes(length, 0xff),
|
|
||||||
"alternating": pattern(length, []byte{0xa5, 0x5a}),
|
|
||||||
"ascending": ascending(length),
|
|
||||||
}
|
|
||||||
for name, buf := range patterns {
|
|
||||||
for _, seed := range []uint16{0, 0xffff, 0x8000} {
|
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
|
||||||
got := Checksum(buf, seed)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("%s len=%d seed=%#x: got %#04x want %#04x",
|
|
||||||
name, length, seed, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func bytes(n int, v byte) []byte {
|
|
||||||
b := make([]byte, n)
|
|
||||||
for i := range b {
|
|
||||||
b[i] = v
|
|
||||||
}
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
func pattern(n int, p []byte) []byte {
|
|
||||||
b := make([]byte, n)
|
|
||||||
for i := range b {
|
|
||||||
b[i] = p[i%len(p)]
|
|
||||||
}
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
func ascending(n int) []byte {
|
|
||||||
b := make([]byte, n)
|
|
||||||
for i := range b {
|
|
||||||
b[i] = byte(i)
|
|
||||||
}
|
|
||||||
return b
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestChecksumTailPaths targets every combination of (SIMD body iterations,
|
|
||||||
// trailing tail bytes) the asm handlers walk through. The tail handlers
|
|
||||||
// peel off 8 → 4 → 2 → 1 byte chunks in turn; this test exercises each by
|
|
||||||
// constructing lengths of the form 64*k + tail for tail ∈ [0, 63] and a
|
|
||||||
// representative spread of k values, including k=0 (no main loop, all tail)
|
|
||||||
// and k=1 (one main loop iter, then tail). It's explicit coverage for
|
|
||||||
// payload sizes that are odd, not divisible by 4, by 8, or by 32.
|
|
||||||
func TestChecksumTailPaths(t *testing.T) {
|
|
||||||
rng := rand.New(rand.NewPCG(42, 17))
|
|
||||||
const padFront = 16
|
|
||||||
const maxK = 8
|
|
||||||
|
|
||||||
pool := make([]byte, 64*maxK+padFront+64)
|
|
||||||
for i := range pool {
|
|
||||||
pool[i] = byte(rng.Uint32())
|
|
||||||
}
|
|
||||||
|
|
||||||
seeds := []uint16{0, 0xffff, 0xabcd}
|
|
||||||
offsets := []int{0, 1, 3, 7, 15} // mix of aligned and odd starts
|
|
||||||
|
|
||||||
for k := 0; k <= maxK; k++ {
|
|
||||||
for tail := 0; tail < 64; tail++ {
|
|
||||||
length := 64*k + tail
|
|
||||||
for _, seed := range seeds {
|
|
||||||
for _, off := range offsets {
|
|
||||||
if off+length > len(pool) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
buf := pool[off : off+length]
|
|
||||||
want := gvisorchecksum.Checksum(buf, seed)
|
|
||||||
got := Checksum(buf, seed)
|
|
||||||
if got != want {
|
|
||||||
t.Fatalf("k=%d tail=%d (len=%d) off=%d seed=%#x: got %#04x want %#04x",
|
|
||||||
k, tail, length, off, seed, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkChecksumTailSizes covers payload sizes that aren't clean multiples
|
|
||||||
// of the SIMD body's 32-byte (amd64) or 16-byte (arm64) chunks, so the tail
|
|
||||||
// handler is meaningfully on the hot path. Sizes are picked to either exercise
|
|
||||||
// every tail branch (tiny lengths) or sit slightly off realistic packet
|
|
||||||
// boundaries (e.g. 1499 = MTU − 1).
|
|
||||||
func BenchmarkChecksumTailSizes(b *testing.B) {
|
|
||||||
sizes := []int{
|
|
||||||
1, 3, 7, 15, 31, // sub-SIMD; entire work is scalar tail
|
|
||||||
33, 35, 47, 63, // one loop32 + assorted tails
|
|
||||||
65, 95, 127, // one loop64 + assorted tails
|
|
||||||
1447, 1471, 1499, 1501, // around MTU
|
|
||||||
8191, 8193, // around USO
|
|
||||||
65531, 65533, // near the kernel max
|
|
||||||
}
|
|
||||||
for _, size := range sizes {
|
|
||||||
buf := make([]byte, size)
|
|
||||||
for i := range buf {
|
|
||||||
buf[i] = byte(i)
|
|
||||||
}
|
|
||||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = gvisorchecksum.Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkChecksum compares the local Checksum to gvisor's at sizes that
|
|
||||||
// match real traffic: a TCP/IP header (60), a typical MSS (1448), a typical
|
|
||||||
// USO size (8192), and the kernel's max GSO superpacket (65535).
|
|
||||||
func BenchmarkChecksum(b *testing.B) {
|
|
||||||
for _, size := range []int{60, 1448, 8192, 65535} {
|
|
||||||
buf := make([]byte, size)
|
|
||||||
for i := range buf {
|
|
||||||
buf[i] = byte(i)
|
|
||||||
}
|
|
||||||
b.Run(fmt.Sprintf("size=%d/local", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
b.Run(fmt.Sprintf("size=%d/gvisor", size), func(b *testing.B) {
|
|
||||||
b.SetBytes(int64(size))
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
_ = gvisorchecksum.Checksum(buf, 0)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+1
-1
@@ -18,7 +18,7 @@ type Device interface {
|
|||||||
Networks() []netip.Prefix
|
Networks() []netip.Prefix
|
||||||
Name() string
|
Name() string
|
||||||
RoutesFor(netip.Addr) routing.Gateways
|
RoutesFor(netip.Addr) routing.Gateways
|
||||||
SupportsMultiqueue() bool
|
SupportsMultiqueue() bool //todo remove?
|
||||||
NewMultiQueueReader() error
|
NewMultiQueueReader() error
|
||||||
Readers() []tio.Queue
|
Readers() []tio.Queue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,358 +0,0 @@
|
|||||||
//go:build !e2e_testing
|
|
||||||
// +build !e2e_testing
|
|
||||||
|
|
||||||
package overlay
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
|
||||||
"runtime"
|
|
||||||
"strings"
|
|
||||||
"syscall"
|
|
||||||
"time"
|
|
||||||
"unsafe"
|
|
||||||
|
|
||||||
"golang.org/x/sys/windows"
|
|
||||||
)
|
|
||||||
|
|
||||||
// networkCategory mirrors NLM_NETWORK_CATEGORY from netlistmgr.h.
|
|
||||||
type networkCategory int32
|
|
||||||
|
|
||||||
const (
|
|
||||||
networkCategoryPublic networkCategory = 0
|
|
||||||
networkCategoryPrivate networkCategory = 1
|
|
||||||
networkCategoryDomainAuthenticated networkCategory = 2
|
|
||||||
)
|
|
||||||
|
|
||||||
func (c networkCategory) String() string {
|
|
||||||
switch c {
|
|
||||||
case networkCategoryPublic:
|
|
||||||
return "public"
|
|
||||||
case networkCategoryPrivate:
|
|
||||||
return "private"
|
|
||||||
case networkCategoryDomainAuthenticated:
|
|
||||||
return "domain"
|
|
||||||
}
|
|
||||||
return fmt.Sprintf("unknown(%d)", c)
|
|
||||||
}
|
|
||||||
|
|
||||||
// parseNetworkCategory accepts the user-supplied tun.network_category. A
|
|
||||||
// second return of false means "leave the category alone".
|
|
||||||
func parseNetworkCategory(s string) (networkCategory, bool, error) {
|
|
||||||
switch strings.ToLower(strings.TrimSpace(s)) {
|
|
||||||
case "", "unset":
|
|
||||||
return 0, false, nil
|
|
||||||
case "public":
|
|
||||||
return networkCategoryPublic, true, nil
|
|
||||||
case "private":
|
|
||||||
return networkCategoryPrivate, true, nil
|
|
||||||
case "domain", "domainauthenticated":
|
|
||||||
return networkCategoryDomainAuthenticated, true, nil
|
|
||||||
}
|
|
||||||
return 0, false, fmt.Errorf("unknown tun.network_category %q (expected public, private, domain, or unset)", s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// CLSID_NetworkListManager {DCB00C01-570F-4A9B-8D69-199FDBA5723B}
|
|
||||||
var clsidNetworkListManager = windows.GUID{
|
|
||||||
Data1: 0xDCB00C01, Data2: 0x570F, Data3: 0x4A9B,
|
|
||||||
Data4: [8]byte{0x8D, 0x69, 0x19, 0x9F, 0xDB, 0xA5, 0x72, 0x3B},
|
|
||||||
}
|
|
||||||
|
|
||||||
// IID_INetworkListManager {DCB00000-570F-4A9B-8D69-199FDBA5723B}
|
|
||||||
var iidINetworkListManager = windows.GUID{
|
|
||||||
Data1: 0xDCB00000, Data2: 0x570F, Data3: 0x4A9B,
|
|
||||||
Data4: [8]byte{0x8D, 0x69, 0x19, 0x9F, 0xDB, 0xA5, 0x72, 0x3B},
|
|
||||||
}
|
|
||||||
|
|
||||||
// x/sys/windows doesn't expose CoCreateInstance, so we bind it ourselves.
|
|
||||||
var procCoCreateInstance = windows.NewLazySystemDLL("ole32.dll").NewProc("CoCreateInstance")
|
|
||||||
|
|
||||||
const clsCtxAll = windows.CLSCTX_INPROC_SERVER | windows.CLSCTX_INPROC_HANDLER |
|
|
||||||
windows.CLSCTX_LOCAL_SERVER | windows.CLSCTX_REMOTE_SERVER
|
|
||||||
|
|
||||||
const (
|
|
||||||
hrSFALSE = 0x00000001
|
|
||||||
hrRPCEChangedMode = 0x80010106
|
|
||||||
)
|
|
||||||
|
|
||||||
type hresult uint32
|
|
||||||
|
|
||||||
func (h hresult) failed() bool { return int32(h) < 0 }
|
|
||||||
func (h hresult) String() string {
|
|
||||||
return fmt.Sprintf("HRESULT 0x%08x", uint32(h))
|
|
||||||
}
|
|
||||||
|
|
||||||
var errAdapterNotFound = errors.New("adapter not present in network connections enumeration")
|
|
||||||
|
|
||||||
// Vtable layouts. Slot order must match the declaration order in netlistmgr.h.
|
|
||||||
// All NLM interfaces here derive from IDispatch, which derives from IUnknown.
|
|
||||||
|
|
||||||
type iUnknownVtbl struct {
|
|
||||||
QueryInterface uintptr
|
|
||||||
AddRef uintptr
|
|
||||||
Release uintptr
|
|
||||||
}
|
|
||||||
|
|
||||||
type iDispatchVtbl struct {
|
|
||||||
iUnknownVtbl
|
|
||||||
GetTypeInfoCount uintptr
|
|
||||||
GetTypeInfo uintptr
|
|
||||||
GetIDsOfNames uintptr
|
|
||||||
Invoke uintptr
|
|
||||||
}
|
|
||||||
|
|
||||||
type iNetworkListManagerVtbl struct {
|
|
||||||
iDispatchVtbl
|
|
||||||
GetNetworks uintptr
|
|
||||||
GetNetwork uintptr
|
|
||||||
GetNetworkConnections uintptr
|
|
||||||
GetNetworkConnection uintptr
|
|
||||||
IsConnectedToInternet uintptr
|
|
||||||
IsConnected uintptr
|
|
||||||
GetConnectivity uintptr
|
|
||||||
}
|
|
||||||
|
|
||||||
type iNetworkListManager struct{ Vtbl *iNetworkListManagerVtbl }
|
|
||||||
|
|
||||||
func (n *iNetworkListManager) Release() {
|
|
||||||
syscall.SyscallN(n.Vtbl.Release, uintptr(unsafe.Pointer(n)))
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n *iNetworkListManager) GetNetworkConnections() (*iEnumNetworkConnections, error) {
|
|
||||||
var enum *iEnumNetworkConnections
|
|
||||||
r1, _, _ := syscall.SyscallN(n.Vtbl.GetNetworkConnections,
|
|
||||||
uintptr(unsafe.Pointer(n)), uintptr(unsafe.Pointer(&enum)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return nil, fmt.Errorf("INetworkListManager.GetNetworkConnections: %s", hr)
|
|
||||||
}
|
|
||||||
return enum, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type iEnumNetworkConnectionsVtbl struct {
|
|
||||||
iDispatchVtbl
|
|
||||||
NewEnum uintptr
|
|
||||||
Next uintptr
|
|
||||||
Skip uintptr
|
|
||||||
Reset uintptr
|
|
||||||
Clone uintptr
|
|
||||||
}
|
|
||||||
|
|
||||||
type iEnumNetworkConnections struct{ Vtbl *iEnumNetworkConnectionsVtbl }
|
|
||||||
|
|
||||||
func (e *iEnumNetworkConnections) Release() {
|
|
||||||
syscall.SyscallN(e.Vtbl.Release, uintptr(unsafe.Pointer(e)))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Next returns the next connection, or (nil, nil) at the end of the enumeration.
|
|
||||||
func (e *iEnumNetworkConnections) Next() (*iNetworkConnection, error) {
|
|
||||||
var conn *iNetworkConnection
|
|
||||||
var fetched uint32
|
|
||||||
r1, _, _ := syscall.SyscallN(e.Vtbl.Next,
|
|
||||||
uintptr(unsafe.Pointer(e)), 1,
|
|
||||||
uintptr(unsafe.Pointer(&conn)), uintptr(unsafe.Pointer(&fetched)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return nil, fmt.Errorf("IEnumNetworkConnections.Next: %s", hr)
|
|
||||||
}
|
|
||||||
if fetched == 0 {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
return conn, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type iNetworkConnectionVtbl struct {
|
|
||||||
iDispatchVtbl
|
|
||||||
GetNetwork uintptr
|
|
||||||
IsConnectedToInternet uintptr
|
|
||||||
IsConnected uintptr
|
|
||||||
GetConnectivity uintptr
|
|
||||||
GetConnectionId uintptr
|
|
||||||
GetAdapterId uintptr
|
|
||||||
GetDomainType uintptr
|
|
||||||
}
|
|
||||||
|
|
||||||
type iNetworkConnection struct{ Vtbl *iNetworkConnectionVtbl }
|
|
||||||
|
|
||||||
func (c *iNetworkConnection) Release() {
|
|
||||||
syscall.SyscallN(c.Vtbl.Release, uintptr(unsafe.Pointer(c)))
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *iNetworkConnection) GetAdapterId() (windows.GUID, error) {
|
|
||||||
var g windows.GUID
|
|
||||||
r1, _, _ := syscall.SyscallN(c.Vtbl.GetAdapterId,
|
|
||||||
uintptr(unsafe.Pointer(c)), uintptr(unsafe.Pointer(&g)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return windows.GUID{}, fmt.Errorf("INetworkConnection.GetAdapterId: %s", hr)
|
|
||||||
}
|
|
||||||
return g, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *iNetworkConnection) GetNetwork() (*iNetwork, error) {
|
|
||||||
var net *iNetwork
|
|
||||||
r1, _, _ := syscall.SyscallN(c.Vtbl.GetNetwork,
|
|
||||||
uintptr(unsafe.Pointer(c)), uintptr(unsafe.Pointer(&net)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return nil, fmt.Errorf("INetworkConnection.GetNetwork: %s", hr)
|
|
||||||
}
|
|
||||||
return net, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type iNetworkVtbl struct {
|
|
||||||
iDispatchVtbl
|
|
||||||
GetName uintptr
|
|
||||||
SetName uintptr
|
|
||||||
GetDescription uintptr
|
|
||||||
SetDescription uintptr
|
|
||||||
GetNetworkId uintptr
|
|
||||||
GetDomainType uintptr
|
|
||||||
GetNetworkConnections uintptr
|
|
||||||
GetTimeCreatedAndConnected uintptr
|
|
||||||
IsConnectedToInternet uintptr
|
|
||||||
IsConnected uintptr
|
|
||||||
GetConnectivity uintptr
|
|
||||||
GetCategory uintptr
|
|
||||||
SetCategory uintptr
|
|
||||||
}
|
|
||||||
|
|
||||||
type iNetwork struct{ Vtbl *iNetworkVtbl }
|
|
||||||
|
|
||||||
func (n *iNetwork) Release() {
|
|
||||||
syscall.SyscallN(n.Vtbl.Release, uintptr(unsafe.Pointer(n)))
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n *iNetwork) GetCategory() (networkCategory, error) {
|
|
||||||
var c networkCategory
|
|
||||||
r1, _, _ := syscall.SyscallN(n.Vtbl.GetCategory,
|
|
||||||
uintptr(unsafe.Pointer(n)), uintptr(unsafe.Pointer(&c)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return 0, fmt.Errorf("INetwork.GetCategory: %s", hr)
|
|
||||||
}
|
|
||||||
return c, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n *iNetwork) SetCategory(c networkCategory) error {
|
|
||||||
r1, _, _ := syscall.SyscallN(n.Vtbl.SetCategory,
|
|
||||||
uintptr(unsafe.Pointer(n)), uintptr(int32(c)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return fmt.Errorf("INetwork.SetCategory: %s", hr)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// coInit initializes COM for the current OS thread. The returned function must
|
|
||||||
// be deferred to balance a successful init. RPC_E_CHANGED_MODE means COM is
|
|
||||||
// already initialized in a different mode on this thread, which is still fine
|
|
||||||
// for our calls but we must not Uninitialize in that case.
|
|
||||||
func coInit() (func(), error) {
|
|
||||||
err := windows.CoInitializeEx(0, windows.COINIT_MULTITHREADED)
|
|
||||||
if err == nil {
|
|
||||||
return windows.CoUninitialize, nil
|
|
||||||
}
|
|
||||||
if e, ok := err.(syscall.Errno); ok {
|
|
||||||
switch uint32(e) {
|
|
||||||
case hrSFALSE:
|
|
||||||
return windows.CoUninitialize, nil
|
|
||||||
case hrRPCEChangedMode:
|
|
||||||
return func() {}, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil, fmt.Errorf("CoInitializeEx: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func createNetworkListManager() (*iNetworkListManager, error) {
|
|
||||||
var nlm *iNetworkListManager
|
|
||||||
r1, _, _ := procCoCreateInstance.Call(
|
|
||||||
uintptr(unsafe.Pointer(&clsidNetworkListManager)),
|
|
||||||
0,
|
|
||||||
uintptr(clsCtxAll),
|
|
||||||
uintptr(unsafe.Pointer(&iidINetworkListManager)),
|
|
||||||
uintptr(unsafe.Pointer(&nlm)),
|
|
||||||
)
|
|
||||||
if hr := hresult(r1); hr.failed() {
|
|
||||||
return nil, fmt.Errorf("CoCreateInstance(NetworkListManager): %s", hr)
|
|
||||||
}
|
|
||||||
return nlm, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// setNetworkCategory locates the network connection bound to adapterGUID and
|
|
||||||
// sets the category of its parent network. Returns errAdapterNotFound if the
|
|
||||||
// adapter is not yet visible in the NLM enumeration.
|
|
||||||
func setNetworkCategory(adapterGUID windows.GUID, cat networkCategory) error {
|
|
||||||
deinit, err := coInit()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
defer deinit()
|
|
||||||
|
|
||||||
nlm, err := createNetworkListManager()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
defer nlm.Release()
|
|
||||||
|
|
||||||
enum, err := nlm.GetNetworkConnections()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
defer enum.Release()
|
|
||||||
|
|
||||||
for {
|
|
||||||
conn, err := enum.Next()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if conn == nil {
|
|
||||||
return errAdapterNotFound
|
|
||||||
}
|
|
||||||
|
|
||||||
guid, err := conn.GetAdapterId()
|
|
||||||
if err != nil || guid != adapterGUID {
|
|
||||||
conn.Release()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
net, err := conn.GetNetwork()
|
|
||||||
conn.Release()
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
err = net.SetCategory(cat)
|
|
||||||
net.Release()
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// applyNetworkCategory polls until the wintun adapter shows up in the NLM
|
|
||||||
// enumeration, then sets the category. Intended to run in its own goroutine.
|
|
||||||
func applyNetworkCategory(l *slog.Logger, adapterGUID windows.GUID, cat networkCategory) {
|
|
||||||
// COM Init/Uninit must be paired on the same OS thread.
|
|
||||||
runtime.LockOSThread()
|
|
||||||
defer runtime.UnlockOSThread()
|
|
||||||
|
|
||||||
const (
|
|
||||||
attempts = 30
|
|
||||||
interval = 500 * time.Millisecond
|
|
||||||
)
|
|
||||||
for i := 0; i < attempts; i++ {
|
|
||||||
err := setNetworkCategory(adapterGUID, cat)
|
|
||||||
if err == nil {
|
|
||||||
l.Info("Set Windows network category", "category", cat.String())
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if !errors.Is(err, errAdapterNotFound) {
|
|
||||||
l.Warn("Failed to set Windows network category", "error", err, "category", cat.String())
|
|
||||||
return
|
|
||||||
}
|
|
||||||
time.Sleep(interval)
|
|
||||||
}
|
|
||||||
l.Warn("Gave up waiting for adapter to appear in NLM enumeration; network category not set",
|
|
||||||
"category", cat.String(),
|
|
||||||
"waited", time.Duration(attempts)*interval,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
@@ -1,109 +0,0 @@
|
|||||||
//go:build !e2e_testing
|
|
||||||
// +build !e2e_testing
|
|
||||||
|
|
||||||
package overlay
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func Test_parseNetworkCategory(t *testing.T) {
|
|
||||||
cases := []struct {
|
|
||||||
in string
|
|
||||||
wantCat networkCategory
|
|
||||||
wantApply bool
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{"", 0, false, false},
|
|
||||||
{"unset", 0, false, false},
|
|
||||||
{" UNSET ", 0, false, false},
|
|
||||||
{"private", networkCategoryPrivate, true, false},
|
|
||||||
{"Private", networkCategoryPrivate, true, false},
|
|
||||||
{" PRIVATE ", networkCategoryPrivate, true, false},
|
|
||||||
{"public", networkCategoryPublic, true, false},
|
|
||||||
{"PUBLIC", networkCategoryPublic, true, false},
|
|
||||||
{"domain", networkCategoryDomainAuthenticated, true, false},
|
|
||||||
{"DomainAuthenticated", networkCategoryDomainAuthenticated, true, false},
|
|
||||||
{"garbage", 0, false, true},
|
|
||||||
{"privates", 0, false, true},
|
|
||||||
}
|
|
||||||
for _, tc := range cases {
|
|
||||||
cat, apply, err := parseNetworkCategory(tc.in)
|
|
||||||
if (err != nil) != tc.wantErr {
|
|
||||||
t.Errorf("parseNetworkCategory(%q) err=%v, wantErr=%v", tc.in, err, tc.wantErr)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if cat != tc.wantCat || apply != tc.wantApply {
|
|
||||||
t.Errorf("parseNetworkCategory(%q) = (%v, %v), want (%v, %v)", tc.in, cat, apply, tc.wantCat, tc.wantApply)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Test_NLM_round_trip exercises every COM call path used by setNetworkCategory
|
|
||||||
// without mutating the host's network state. It validates the CLSID/IID
|
|
||||||
// constants and every vtable index by enumerating connections, fetching the
|
|
||||||
// adapter id and parent network, reading the current category, and writing it
|
|
||||||
// back unchanged.
|
|
||||||
//
|
|
||||||
// Requires Windows but does not require admin or the wintun driver. Skips if
|
|
||||||
// no network connections are available (unlikely outside of an isolated
|
|
||||||
// container).
|
|
||||||
func Test_NLM_round_trip(t *testing.T) {
|
|
||||||
deinit, err := coInit()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("coInit: %v", err)
|
|
||||||
}
|
|
||||||
defer deinit()
|
|
||||||
|
|
||||||
nlm, err := createNetworkListManager()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("createNetworkListManager: %v", err)
|
|
||||||
}
|
|
||||||
defer nlm.Release()
|
|
||||||
|
|
||||||
enum, err := nlm.GetNetworkConnections()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("GetNetworkConnections: %v", err)
|
|
||||||
}
|
|
||||||
defer enum.Release()
|
|
||||||
|
|
||||||
saw := 0
|
|
||||||
for {
|
|
||||||
conn, err := enum.Next()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("EnumNetworkConnections.Next: %v", err)
|
|
||||||
}
|
|
||||||
if conn == nil {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
saw++
|
|
||||||
|
|
||||||
if _, err := conn.GetAdapterId(); err != nil {
|
|
||||||
conn.Release()
|
|
||||||
t.Fatalf("INetworkConnection.GetAdapterId: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
net, err := conn.GetNetwork()
|
|
||||||
conn.Release()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("INetworkConnection.GetNetwork: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
cat, err := net.GetCategory()
|
|
||||||
if err != nil {
|
|
||||||
net.Release()
|
|
||||||
t.Fatalf("INetwork.GetCategory: %v", err)
|
|
||||||
}
|
|
||||||
// Set to the current value so the host's NLM state is unchanged but
|
|
||||||
// SetCategory's vtable slot is still validated end-to-end.
|
|
||||||
if err := net.SetCategory(cat); err != nil {
|
|
||||||
net.Release()
|
|
||||||
t.Fatalf("INetwork.SetCategory(%v): %v", cat, err)
|
|
||||||
}
|
|
||||||
net.Release()
|
|
||||||
}
|
|
||||||
|
|
||||||
if saw == 0 {
|
|
||||||
t.Skip("no NLM network connections available; skipping round-trip")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -31,7 +31,7 @@ func (NoopTun) Name() string {
|
|||||||
return "noop"
|
return "noop"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Read() ([]tio.Packet, error) {
|
func (NoopTun) Read() ([][]byte, error) {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -8,43 +8,34 @@ import (
|
|||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
type offloadQueueSet struct {
|
type offloadContainer struct {
|
||||||
pq []*Offload
|
pq []*Offload
|
||||||
// pqi is exactly the same as pq, but stored as the interface type
|
// pqi is exactly the same as pq, but stored as the interface type
|
||||||
pqi []Queue
|
pqi []Queue
|
||||||
shutdownFd int
|
shutdownFd int
|
||||||
// usoEnabled is true when newTun successfully negotiated TUN_F_USO4|6
|
|
||||||
// with the kernel. Queues created by Add inherit this and surface it
|
|
||||||
// via Offload.USOSupported so coalescers can gate USO emission.
|
|
||||||
usoEnabled bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewOffloadQueueSet creates a QueueSet that uses virtio_net_hdr to do
|
func NewOffloadContainer() (Container, error) {
|
||||||
// TSO segmentation in userspace. usoEnabled tells downstream queues whether
|
|
||||||
// the kernel agreed to deliver/accept GSO_UDP_L4 superpackets — coalescers
|
|
||||||
// should fall back to per-packet writes when this is false.
|
|
||||||
func NewOffloadQueueSet(usoEnabled bool) (QueueSet, error) {
|
|
||||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
out := &offloadQueueSet{
|
out := &offloadContainer{
|
||||||
pq: []*Offload{},
|
pq: []*Offload{},
|
||||||
pqi: []Queue{},
|
pqi: []Queue{},
|
||||||
shutdownFd: shutdownFd,
|
shutdownFd: shutdownFd,
|
||||||
usoEnabled: usoEnabled,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *offloadQueueSet) Queues() []Queue {
|
func (c *offloadContainer) Queues() []Queue {
|
||||||
return c.pqi
|
return c.pqi
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *offloadQueueSet) Add(fd int) error {
|
func (c *offloadContainer) Add(fd int) error {
|
||||||
x, err := newOffload(fd, c.shutdownFd, c.usoEnabled)
|
x, err := newOffload(fd, c.shutdownFd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -54,14 +45,14 @@ func (c *offloadQueueSet) Add(fd int) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *offloadQueueSet) wakeForShutdown() error {
|
func (c *offloadContainer) wakeForShutdown() error {
|
||||||
var buf [8]byte
|
var buf [8]byte
|
||||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||||
_, err := unix.Write(c.shutdownFd, buf[:])
|
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *offloadQueueSet) Close() error {
|
func (c *offloadContainer) Close() error {
|
||||||
errs := []error{}
|
errs := []error{}
|
||||||
|
|
||||||
// Signal all readers blocked in poll to wake up and exit
|
// Signal all readers blocked in poll to wake up and exit
|
||||||
@@ -8,20 +8,20 @@ import (
|
|||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
type pollQueueSet struct {
|
type pollContainer struct {
|
||||||
pq []*Poll
|
pq []*Poll
|
||||||
// pqi is exactly the same as pq, but stored as the interface type
|
// pqi is exactly the same as pq, but stored as the interface type
|
||||||
pqi []Queue
|
pqi []Queue
|
||||||
shutdownFd int
|
shutdownFd int
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewPollQueueSet() (QueueSet, error) {
|
func NewPollContainer() (Container, error) {
|
||||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
out := &pollQueueSet{
|
out := &pollContainer{
|
||||||
pq: []*Poll{},
|
pq: []*Poll{},
|
||||||
pqi: []Queue{},
|
pqi: []Queue{},
|
||||||
shutdownFd: shutdownFd,
|
shutdownFd: shutdownFd,
|
||||||
@@ -30,11 +30,11 @@ func NewPollQueueSet() (QueueSet, error) {
|
|||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *pollQueueSet) Queues() []Queue {
|
func (c *pollContainer) Queues() []Queue {
|
||||||
return c.pqi
|
return c.pqi
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *pollQueueSet) Add(fd int) error {
|
func (c *pollContainer) Add(fd int) error {
|
||||||
x, err := newPoll(fd, c.shutdownFd)
|
x, err := newPoll(fd, c.shutdownFd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -45,14 +45,14 @@ func (c *pollQueueSet) Add(fd int) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *pollQueueSet) wakeForShutdown() error {
|
func (c *pollContainer) wakeForShutdown() error {
|
||||||
var buf [8]byte
|
var buf [8]byte
|
||||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||||
_, err := unix.Write(int(c.shutdownFd), buf[:])
|
_, err := unix.Write(int(c.shutdownFd), buf[:])
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *pollQueueSet) Close() error {
|
func (c *pollContainer) Close() error {
|
||||||
errs := []error{}
|
errs := []error{}
|
||||||
|
|
||||||
if err := c.wakeForShutdown(); err != nil {
|
if err := c.wakeForShutdown(); err != nil {
|
||||||
@@ -1,65 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
|
|
||||||
package tio
|
|
||||||
|
|
||||||
import "testing"
|
|
||||||
|
|
||||||
// fakeBatch stands in for batch.TxBatcher inside the bench — same shape
|
|
||||||
// of pointer-capturing closure that sendInsideMessage builds.
|
|
||||||
type fakeBatch struct{ buf [65536]byte }
|
|
||||||
|
|
||||||
func (b *fakeBatch) Reserve(sz int) []byte { return b.buf[:sz] }
|
|
||||||
func (b *fakeBatch) Commit([]byte) {}
|
|
||||||
|
|
||||||
type fakeHostInfo struct {
|
|
||||||
remoteIndexId uint32
|
|
||||||
counter uint64
|
|
||||||
}
|
|
||||||
type fakeIface struct {
|
|
||||||
rebindCount uint8
|
|
||||||
hi *fakeHostInfo
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkSegmentSuperpacketAllocsTSO measures allocation per
|
|
||||||
// SegmentSuperpacket call when a closure captures pointer-bearing
|
|
||||||
// receivers — the realistic shape of sendInsideMessage's closure.
|
|
||||||
func BenchmarkSegmentSuperpacketAllocsTSO(b *testing.B) {
|
|
||||||
const mss = 1400
|
|
||||||
const numSeg = 32
|
|
||||||
pkt := buildTSOv6(mss*numSeg, mss)
|
|
||||||
gso := GSOInfo{
|
|
||||||
Size: mss,
|
|
||||||
HdrLen: 60, // 40 (IPv6) + 20 (TCP)
|
|
||||||
CsumStart: 40,
|
|
||||||
Proto: GSOProtoTCP,
|
|
||||||
}
|
|
||||||
p := Packet{Bytes: pkt, GSO: gso}
|
|
||||||
|
|
||||||
hi := &fakeHostInfo{remoteIndexId: 0xdeadbeef}
|
|
||||||
f := &fakeIface{rebindCount: 7, hi: hi}
|
|
||||||
fb := &fakeBatch{}
|
|
||||||
|
|
||||||
// SegmentSuperpacket consumes pkt destructively; refresh from a master
|
|
||||||
// copy each iter (matches the production pattern where every TUN read
|
|
||||||
// hands the segmenter a fresh kernel-supplied buffer).
|
|
||||||
master := append([]byte(nil), pkt...)
|
|
||||||
work := make([]byte, len(pkt))
|
|
||||||
p.Bytes = work
|
|
||||||
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
copy(work, master)
|
|
||||||
err := SegmentSuperpacket(p, func(seg []byte) error {
|
|
||||||
out := fb.Reserve(16 + len(seg) + 16)
|
|
||||||
out[0] = byte(f.rebindCount)
|
|
||||||
out[1] = byte(hi.counter)
|
|
||||||
hi.counter++
|
|
||||||
fb.Commit(out)
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
b.Fatalf("SegmentSuperpacket: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,22 +0,0 @@
|
|||||||
//go:build !linux || android || e2e_testing
|
|
||||||
|
|
||||||
package tio
|
|
||||||
|
|
||||||
import "fmt"
|
|
||||||
|
|
||||||
func protoFromGSOType(_ uint8) (GSOProto, error) {
|
|
||||||
return 0, fmt.Errorf("GSO unsupported")
|
|
||||||
}
|
|
||||||
|
|
||||||
// SegmentSuperpacket invokes fn once per segment of pkt. On non-Linux
|
|
||||||
// builds (and Android/e2e_testing) this package does not provide a Queue
|
|
||||||
// implementation, so any caller that does construct a Packet here can only
|
|
||||||
// be operating on non-superpacket bytes and the stub forwards them
|
|
||||||
// directly. A non-zero GSO field is a programming error from the caller
|
|
||||||
// and returns an explicit error rather than silently misbehaving.
|
|
||||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
|
||||||
if pkt.GSO.IsSuperpacket() {
|
|
||||||
return fmt.Errorf("tio: GSO superpacket on platform without segmentation support")
|
|
||||||
}
|
|
||||||
return fn(pkt.Bytes)
|
|
||||||
}
|
|
||||||
+36
-139
@@ -4,167 +4,64 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
)
|
)
|
||||||
|
|
||||||
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
|
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||||
type QueueSet interface {
|
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||||
|
const defaultBatchBufSize = 65535
|
||||||
|
|
||||||
|
// Container holds one or many Queue objects and helps close them in an orderly way
|
||||||
|
type Container interface {
|
||||||
io.Closer
|
io.Closer
|
||||||
Queues() []Queue
|
Queues() []Queue
|
||||||
|
|
||||||
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
|
// Add takes a tun fd, adds it to the container, and prepares it for use as a Queue
|
||||||
Add(fd int) error
|
Add(fd int) error
|
||||||
}
|
|
||||||
|
|
||||||
// Capabilities advertises which kernel offload features a Queue
|
io.Closer
|
||||||
// successfully negotiated. Callers consult this to decide which coalescers
|
|
||||||
// to wire onto the write path — a Queue without TSO can't usefully accept a
|
|
||||||
// TCPCoalescer, and a Queue without USO can't accept a UDPCoalescer.
|
|
||||||
type Capabilities struct {
|
|
||||||
// TSO means the FD was opened with IFF_VNET_HDR and the kernel agreed
|
|
||||||
// to TUN_F_TSO4|TSO6 — i.e. WriteGSO with GSOProtoTCP is safe.
|
|
||||||
TSO bool
|
|
||||||
// USO means the kernel additionally agreed to TUN_F_USO4|USO6, so
|
|
||||||
// WriteGSO with GSOProtoUDP is safe. Linux ≥ 6.2.
|
|
||||||
USO bool
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Queue is a readable/writable Poll queue. One Queue is driven by a single
|
// Queue is a readable/writable Poll queue. One Queue is driven by a single
|
||||||
// read goroutine plus a single writer (see Write below).
|
// read goroutine plus concurrent writers (see Write / WriteReject below).
|
||||||
type Queue interface {
|
type Queue interface {
|
||||||
io.Closer
|
io.Closer
|
||||||
|
|
||||||
// Read returns one or more packets. The returned Packet.Bytes slices
|
// Read returns one or more packets. The returned slices are borrowed
|
||||||
// are borrowed from the Queue's internal buffer and are only valid
|
// from the Queue's internal buffer and are only valid until the next
|
||||||
// until the next Read or Close on this Queue - callers must encrypt
|
// Read or Close on this Queue - callers must encrypt or copy each
|
||||||
// or copy each slice before the next call. A Packet may carry a
|
// slice before the next call. Not safe for concurrent Reads.
|
||||||
// GSO/USO superpacket (see GSOInfo); when GSO.IsSuperpacket() is
|
Read() ([][]byte, error)
|
||||||
// true the caller must segment Bytes before treating it as a single
|
|
||||||
// IP datagram. Not safe for concurrent Reads.
|
|
||||||
Read() ([]Packet, error)
|
|
||||||
|
|
||||||
// Write emits a single packet on the plaintext (outside→inside)
|
// Write emits a single packet on the plaintext (outside→inside)
|
||||||
// delivery path. Not safe for concurrent Writes.
|
// delivery path. Not safe for concurrent Writes.
|
||||||
Write(p []byte) (int, error)
|
Write(p []byte) (int, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Packet is the unit Queue.Read returns. Bytes points into the queue's
|
// GSOWriter is implemented by Queues that can emit a TCP TSO superpacket
|
||||||
// internal buffer and is only valid until the next Read or Close on the
|
|
||||||
// queue that produced it. GSO is the zero value for an already-segmented
|
|
||||||
// IP datagram; when non-zero it describes a kernel-supplied TSO/USO
|
|
||||||
// superpacket the caller must segment before consuming.
|
|
||||||
type Packet struct {
|
|
||||||
Bytes []byte
|
|
||||||
GSO GSOInfo
|
|
||||||
}
|
|
||||||
|
|
||||||
// GSOInfo describes a kernel-supplied superpacket sitting in Packet.Bytes.
|
|
||||||
// The zero value means "not a superpacket" — Bytes is one regular IP
|
|
||||||
// datagram and no segmentation is required.
|
|
||||||
type GSOInfo struct {
|
|
||||||
// Size is the GSO segment size: max payload bytes per segment
|
|
||||||
// (== TCP MSS for TSO, == UDP payload chunk for USO). Zero means
|
|
||||||
// not a superpacket.
|
|
||||||
Size uint16
|
|
||||||
// HdrLen is the total L3+L4 header length within Bytes (already
|
|
||||||
// corrected via correctHdrLen, so safe to slice on).
|
|
||||||
HdrLen uint16
|
|
||||||
// CsumStart is the L4 header offset inside Bytes (== L3 header
|
|
||||||
// length).
|
|
||||||
CsumStart uint16
|
|
||||||
// Proto picks the L4 protocol (TCP or UDP) so the segmenter knows
|
|
||||||
// which checksum/header layout to apply.
|
|
||||||
Proto GSOProto
|
|
||||||
}
|
|
||||||
|
|
||||||
// IsSuperpacket reports whether g describes a multi-segment GSO/USO
|
|
||||||
// superpacket that needs segmentation before its bytes can be encrypted
|
|
||||||
// and sent on the wire.
|
|
||||||
func (g GSOInfo) IsSuperpacket() bool { return g.Size > 0 }
|
|
||||||
|
|
||||||
// Clone returns a Packet whose Bytes is a freshly allocated copy of p.Bytes,
|
|
||||||
// safe to retain past the next Read or Close on the originating Queue.
|
|
||||||
// GSO metadata is copied verbatim. Use this only when a caller genuinely
|
|
||||||
// needs to outlive the borrowed-slice contract — the hot path reads should
|
|
||||||
// continue to consume the borrow synchronously to avoid the allocation.
|
|
||||||
func (p Packet) Clone() Packet {
|
|
||||||
if p.Bytes == nil {
|
|
||||||
return p
|
|
||||||
}
|
|
||||||
cp := make([]byte, len(p.Bytes))
|
|
||||||
copy(cp, p.Bytes)
|
|
||||||
return Packet{Bytes: cp, GSO: p.GSO}
|
|
||||||
}
|
|
||||||
|
|
||||||
// CapsProvider is an optional interface implemented by Queues that
|
|
||||||
// successfully negotiated kernel offload features at open time. Callers
|
|
||||||
// pick a write-path coalescer based on the result. Queues that don't
|
|
||||||
// implement it are treated as having no offload capability — callers must
|
|
||||||
// fall back to plain per-packet writes.
|
|
||||||
type CapsProvider interface {
|
|
||||||
Capabilities() Capabilities
|
|
||||||
}
|
|
||||||
|
|
||||||
// QueueCapabilities returns q's negotiated offload capabilities, or the
|
|
||||||
// zero value when q does not advertise any.
|
|
||||||
func QueueCapabilities(q Queue) Capabilities {
|
|
||||||
if cp, ok := q.(CapsProvider); ok {
|
|
||||||
return cp.Capabilities()
|
|
||||||
}
|
|
||||||
return Capabilities{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// GSOProto selects the L4 protocol for a GSO superpacket. Determines which
|
|
||||||
// VIRTIO_NET_HDR_GSO_* type the writer stamps and which checksum offset
|
|
||||||
// inside the transport header virtio NEEDS_CSUM expects.
|
|
||||||
type GSOProto uint8
|
|
||||||
|
|
||||||
const (
|
|
||||||
GSOProtoTCP GSOProto = iota
|
|
||||||
GSOProtoUDP
|
|
||||||
)
|
|
||||||
|
|
||||||
// GSOWriter is implemented by Queues that can emit a TCP or UDP superpacket
|
|
||||||
// assembled from a header prefix plus one or more borrowed payload
|
// assembled from a header prefix plus one or more borrowed payload
|
||||||
// fragments, in a single vectored write (writev with a leading
|
// fragments, in a single vectored write (writev with a leading
|
||||||
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
||||||
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
||||||
// support do not implement this interface and coalescing is skipped.
|
// support return false from GSOSupported and coalescing is skipped.
|
||||||
//
|
//
|
||||||
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have
|
// hdr contains the IPv4/IPv6 + TCP header prefix (mutable - callers will
|
||||||
// filled in total length and IP csum). transportHdr is the TCP or UDP
|
// have filled in total length and pseudo-header partial). pays are
|
||||||
// header (mutable - the L4 checksum field must hold the pseudo-header
|
// non-overlapping payload fragments whose concatenation is the full
|
||||||
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
// superpacket payload; they are read-only from the writer's perspective
|
||||||
// pays are non-overlapping payload fragments whose concatenation is the
|
// and must remain valid until the call returns. gsoSize is the MSS:
|
||||||
// full superpacket payload; they are read-only from the writer's
|
// every segment except possibly the last is exactly that many bytes.
|
||||||
// perspective and must remain valid until the call returns. Every segment
|
// csumStart is the byte offset where the TCP header begins within hdr.
|
||||||
// in pays except possibly the last is exactly the same size. proto picks
|
|
||||||
// the L4 protocol so the writer knows which GSOType / CsumOffset to set.
|
|
||||||
//
|
//
|
||||||
// Callers should also consult CapsProvider (via SupportsGSO or
|
// # TODO fold into Queue
|
||||||
// QueueCapabilities) for the per-protocol negotiated capability; an
|
//
|
||||||
// implementation of GSOWriter is necessary but not sufficient since USO
|
// hdr's TCP checksum field must already hold the pseudo-header partial
|
||||||
// may not have been negotiated even when TSO was.
|
// sum (single-fold, not inverted), per virtio NEEDS_CSUM semantics.
|
||||||
type GSOWriter interface {
|
type GSOWriter interface {
|
||||||
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
// WriteGSO emits a TCP TSO superpacket in a single writev. hdr is the
|
||||||
}
|
// IPv4/IPv6 + TCP header prefix (already finalized — total length, IP csum,
|
||||||
|
// and TCP pseudo-header partial set by the caller). pays are payload
|
||||||
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
// fragments whose concatenation forms the full coalesced payload; each
|
||||||
// queue advertises the negotiated capability for `want`. A writer that
|
// slice is read-only and must stay valid until return.
|
||||||
// implements GSOWriter but not CapsProvider is treated as permissive
|
// every segment in pays except possibly the last is exactly the same size.
|
||||||
// (used by tests and fakes that don't negotiate).
|
// csumStart is the byte offset where the TCP header begins within hdr.
|
||||||
func SupportsGSO(w any, want GSOProto) (GSOWriter, bool) {
|
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte) error
|
||||||
gw, ok := w.(GSOWriter)
|
GSOSupported() bool
|
||||||
if !ok {
|
|
||||||
return nil, false
|
|
||||||
}
|
|
||||||
cp, ok := w.(CapsProvider)
|
|
||||||
if !ok {
|
|
||||||
return gw, true
|
|
||||||
}
|
|
||||||
caps := cp.Capabilities()
|
|
||||||
switch want {
|
|
||||||
case GSOProtoTCP:
|
|
||||||
return gw, caps.TSO
|
|
||||||
case GSOProtoUDP:
|
|
||||||
return gw, caps.USO
|
|
||||||
}
|
|
||||||
return gw, false
|
|
||||||
}
|
}
|
||||||
|
|||||||
+94
-202
@@ -3,7 +3,6 @@ package tio
|
|||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"log/slog"
|
|
||||||
"os"
|
"os"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -11,48 +10,38 @@ import (
|
|||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// tunRxBufSize is the per-Read worst-case footprint inside rxBuf: one
|
// Space for segmented output. Worst case is many small segments, each paying
|
||||||
// kernel-supplied packet body, which is at most ~64 KiB (tunReadBufSize).
|
// an IP+TCP header. Should be a multiple of 64KiB.
|
||||||
// Segmentation happens at encrypt time on a per-routine MTU-sized scratch
|
// const tunSegBufSize = 0xffff * 8 TODO larger? config?
|
||||||
// (see SegmentSuperpacket), so rxBuf only holds raw kernel-supplied bytes.
|
const tunSegBufSize = 131072
|
||||||
// We round up to give comfortable margin for the drain headroom check
|
|
||||||
// below.
|
|
||||||
const tunRxBufSize = 64 * 1024
|
|
||||||
|
|
||||||
// tunRxBufCap is the total size we allocate for the per-reader rx
|
// tunSegBufCap is the total size we allocate for the per-reader segment
|
||||||
// buffer. With reads landing directly in rxBuf, each drain iteration
|
// buffer. It is sized as one worst-case TSO superpacket (tunSegBufSize) plus
|
||||||
// consumes up to tunRxBufSize of headroom for the kernel-supplied bytes.
|
// the same again as drain headroom so a Read wake can accumulate
|
||||||
// Sized to two such iterations so the initial blocking read plus one
|
// additional packets after an initial big read without overflowing.
|
||||||
// drain read both fit without partial-drop.
|
const tunSegBufCap = tunSegBufSize * 2
|
||||||
const tunRxBufCap = tunRxBufSize * 2
|
|
||||||
|
|
||||||
// tunDrainCap caps how many packets a single Read will accumulate via
|
// tunDrainCap caps how many packets a single Read will accumulate via
|
||||||
// the post-wake drain loop. Sized to soak up a burst of small ACKs while
|
// the post-wake drain loop. Sized to soak up a burst of small ACKs while
|
||||||
// bounding how much work a single caller holds before handing off.
|
// bounding how much work a single caller holds before handing off.
|
||||||
const tunDrainCap = 64
|
const tunDrainCap = 64 //256
|
||||||
|
|
||||||
// gsoMaxIovs caps the iovec budget WriteGSO assembles per call: 3 fixed
|
// gsoInitialPayIovs is the starting capacity (in payload fragments) of
|
||||||
// entries (virtio_net_hdr, IP hdr, transport hdr) plus up to gsoMaxIovs-3
|
// Offload.gsoIovs. Sized to cover the default coalesce segment cap without
|
||||||
// payload fragments. Sized comfortably above the typical kernel GSO
|
// any reallocations.
|
||||||
// segment cap (Linux UDP_GRO is 64) so realistic coalesced bursts never
|
const gsoInitialPayIovs = 66
|
||||||
// touch the limit. iovecs are tiny (16 bytes), so the entire scratch is
|
|
||||||
// 4 KiB — fine to keep resident on every queue. WriteGSO returns an error
|
|
||||||
// rather than reallocating when a caller exceeds this budget.
|
|
||||||
const gsoMaxIovs = 256
|
|
||||||
|
|
||||||
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN
|
// validVnetHdr is the 10-byte virtio_net_hdr we prepend to every non-GSO TUN
|
||||||
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb
|
// write. Only flag set is VIRTIO_NET_HDR_F_DATA_VALID, which marks the skb
|
||||||
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checksum
|
// CHECKSUM_UNNECESSARY so the receiving network stack skips L4 checks
|
||||||
// verification. All packets that reach the plain Write paths already carry
|
// verification. All packets that reach the plain Write paths
|
||||||
// a valid L4 checksum (either supplied by a remote peer whose ciphertext we
|
// already carry a valid L4 checksum (either supplied by a remote peer whose
|
||||||
// AEAD-authenticated, produced by segmentTCPYield/segmentUDPYield during
|
// ciphertext we AEAD-authenticated, or produced by finishChecksum during TSO
|
||||||
// superpacket segmentation, or built locally by CreateRejectPacket), so
|
// segmentation, or built locally by CreateRejectPacket), so trusting them is
|
||||||
// trusting them is safe.
|
// safe.
|
||||||
var validVnetHdr = [virtio.Size]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
var validVnetHdr = [virtioNetHdrLen]byte{unix.VIRTIO_NET_HDR_F_DATA_VALID}
|
||||||
|
|
||||||
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
// Offload wraps a TUN file descriptor with poll-based reads. The FD provided will be changed to non-blocking.
|
||||||
// A shared eventfd allows Close to wake all readers blocked in poll.
|
// A shared eventfd allows Close to wake all readers blocked in poll.
|
||||||
@@ -61,42 +50,24 @@ type Offload struct {
|
|||||||
shutdownFd int
|
shutdownFd int
|
||||||
readPoll [2]unix.PollFd
|
readPoll [2]unix.PollFd
|
||||||
writePoll [2]unix.PollFd
|
writePoll [2]unix.PollFd
|
||||||
// writeLock serializes blockOnWrite's read+clear of writePoll[*].Revents.
|
writeLock sync.Mutex //there's more than one potential write source per-routine, so we need this to protect writePoll
|
||||||
// Any goroutine that calls Write may end up parked in poll(2); without
|
closed atomic.Bool
|
||||||
// the lock concurrent waiters could race the Revents reset and lose
|
readBuf []byte // scratch for a single raw read (virtio hdr + superpacket)
|
||||||
// events.
|
segBuf []byte // backing store for segmented output
|
||||||
writeLock sync.Mutex
|
segOff int // cursor into segBuf for the current Read drain
|
||||||
closed atomic.Bool
|
pending [][]byte // segments returned from the most recent Read
|
||||||
rxBuf []byte // backing store for kernel-handed packets read this drain
|
|
||||||
rxOff int // cursor into rxBuf for the current Read drain
|
|
||||||
pending []Packet // packets returned from the most recent Read
|
|
||||||
|
|
||||||
// readVnetScratch holds the 10-byte virtio_net_hdr split off the front of
|
|
||||||
// every TUN read via readv(2). Decoupling the header from the packet body
|
|
||||||
// lets us read the body directly into rxBuf at the current rxOff with
|
|
||||||
// no userspace copy on the GSO_NONE fast path.
|
|
||||||
readVnetScratch [virtio.Size]byte
|
|
||||||
// readIovs is the readv(2) iovec scratch wired once at construction —
|
|
||||||
// iovec[0] points at readVnetScratch; iovec[1].Base/Len is updated per
|
|
||||||
// read to address the current rxBuf slot.
|
|
||||||
readIovs [2]unix.Iovec
|
|
||||||
|
|
||||||
// usoEnabled records whether the kernel agreed to TUN_F_USO* on this FD,
|
|
||||||
// so writers can decide whether emitting GSO_UDP_L4 superpackets is safe.
|
|
||||||
usoEnabled bool
|
|
||||||
|
|
||||||
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
// gsoHdrBuf is a per-queue 10-byte scratch for the virtio_net_hdr emitted
|
||||||
// by WriteGSO. Kept separate from the read-only package-level validVnetHdr
|
// by WriteGSO. Separate from validVnetHdr so a concurrent non-GSO Write on
|
||||||
// so non-GSO Writes can ship that constant directly while WriteGSO
|
// another queue never observes a half-written header.
|
||||||
// rewrites this scratch on every call.
|
gsoHdrBuf [virtioNetHdrLen]byte
|
||||||
gsoHdrBuf [virtio.Size]byte
|
// gsoIovs is the writev iovec scratch for WriteGSO. Sized to hold the
|
||||||
// gsoIovs is the writev iovec scratch for WriteGSO. Pre-sized to
|
// virtio header + IP/TCP header + up to gsoInitialPayIovs payload
|
||||||
// gsoMaxIovs at construction; never grown. WriteGSO returns an error
|
// fragments; grown on demand if a coalescer pushes more.
|
||||||
// (and drops the call) if a caller hands it more fragments than fit.
|
|
||||||
gsoIovs []unix.Iovec
|
gsoIovs []unix.Iovec
|
||||||
}
|
}
|
||||||
|
|
||||||
func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
|
func newOffload(fd int, shutdownFd int) (*Offload, error) {
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
||||||
}
|
}
|
||||||
@@ -104,8 +75,8 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
|
|||||||
out := &Offload{
|
out := &Offload{
|
||||||
fd: fd,
|
fd: fd,
|
||||||
shutdownFd: shutdownFd,
|
shutdownFd: shutdownFd,
|
||||||
usoEnabled: usoEnabled,
|
|
||||||
closed: atomic.Bool{},
|
closed: atomic.Bool{},
|
||||||
|
readBuf: make([]byte, virtioNetHdrLen+tunReadBufSize),
|
||||||
readPoll: [2]unix.PollFd{
|
readPoll: [2]unix.PollFd{
|
||||||
{Fd: int32(fd), Events: unix.POLLIN},
|
{Fd: int32(fd), Events: unix.POLLIN},
|
||||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||||
@@ -116,17 +87,12 @@ func newOffload(fd int, shutdownFd int, usoEnabled bool) (*Offload, error) {
|
|||||||
},
|
},
|
||||||
writeLock: sync.Mutex{},
|
writeLock: sync.Mutex{},
|
||||||
|
|
||||||
rxBuf: make([]byte, tunRxBufCap),
|
segBuf: make([]byte, tunSegBufCap),
|
||||||
gsoIovs: make([]unix.Iovec, 2, gsoMaxIovs),
|
gsoIovs: make([]unix.Iovec, 2, 2+gsoInitialPayIovs),
|
||||||
}
|
}
|
||||||
|
|
||||||
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
|
out.gsoIovs[0].Base = &out.gsoHdrBuf[0]
|
||||||
out.gsoIovs[0].SetLen(virtio.Size)
|
out.gsoIovs[0].SetLen(virtioNetHdrLen)
|
||||||
|
|
||||||
// readIovs[0] is wired once to the virtio_net_hdr scratch; per-read we
|
|
||||||
// only repoint readIovs[1] at the next rxBuf slot (see readPacket).
|
|
||||||
out.readIovs[0].Base = &out.readVnetScratch[0]
|
|
||||||
out.readIovs[0].SetLen(virtio.Size)
|
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
@@ -185,66 +151,40 @@ func (r *Offload) blockOnWrite() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// readPacket issues a single readv(2) splitting the virtio_net_hdr off
|
func (r *Offload) readRaw(buf []byte) (int, error) {
|
||||||
// into readVnetScratch and reading the packet body directly into rxBuf at
|
|
||||||
// the current rxOff. Returns the body length (zero virtio header bytes,
|
|
||||||
// just the IP packet/superpacket). block controls whether EAGAIN is
|
|
||||||
// retried via poll: the initial read of a drain blocks; subsequent drain
|
|
||||||
// reads do not.
|
|
||||||
//
|
|
||||||
// The body iovec capacity is always tunReadBufSize; callers (the Read
|
|
||||||
// drain loop) gate entry on tunRxBufCap-rxOff >= tunRxBufSize, sized to
|
|
||||||
// hold one worst-case kernel-supplied packet body. Without that gate the
|
|
||||||
// body iovec could be smaller than the next inbound packet and the
|
|
||||||
// kernel would truncate.
|
|
||||||
func (r *Offload) readPacket(block bool) (int, error) {
|
|
||||||
for {
|
for {
|
||||||
r.readIovs[1].Base = &r.rxBuf[r.rxOff]
|
if n, err := unix.Read(r.fd, buf); err == nil {
|
||||||
r.readIovs[1].SetLen(tunReadBufSize)
|
return n, nil
|
||||||
n, _, errno := syscall.Syscall(unix.SYS_READV, uintptr(r.fd), uintptr(unsafe.Pointer(&r.readIovs[0])), uintptr(len(r.readIovs)))
|
} else if err == unix.EAGAIN {
|
||||||
if errno == 0 {
|
if err = r.blockOnRead(); err != nil {
|
||||||
if int(n) < virtio.Size {
|
|
||||||
return 0, io.ErrShortWrite
|
|
||||||
}
|
|
||||||
return int(n) - virtio.Size, nil
|
|
||||||
}
|
|
||||||
if errno == unix.EAGAIN {
|
|
||||||
if !block {
|
|
||||||
return 0, errno
|
|
||||||
}
|
|
||||||
if err := r.blockOnRead(); err != nil {
|
|
||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
continue
|
continue
|
||||||
}
|
} else if err == unix.EINTR {
|
||||||
if errno == unix.EINTR {
|
|
||||||
continue
|
continue
|
||||||
}
|
} else if err == unix.EBADF {
|
||||||
if errno == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
return 0, os.ErrClosed
|
||||||
|
} else {
|
||||||
|
return 0, err
|
||||||
}
|
}
|
||||||
return 0, errno
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Read returns one or more packets from the tun. Each Packet either
|
// Read reads one or more superpackets from the tun and returns the
|
||||||
// carries a single ready-to-use IP datagram (GSO zero) or a TSO/USO
|
// resulting packets. The first read blocks via poll; once the fd is known
|
||||||
// superpacket plus the GSOInfo a caller needs to segment it (see
|
// readable we drain additional packets non-blocking until the kernel queue
|
||||||
// SegmentSuperpacket). The first read blocks via poll; once the fd is
|
// is empty (EAGAIN), we've collected tunDrainCap packets, or we're out of
|
||||||
// known readable we drain additional packets non-blocking until the
|
// segBuf headroom. This amortizes the poll wake over bursts of small
|
||||||
// kernel queue is empty (EAGAIN), we've collected tunDrainCap packets,
|
// packets (e.g. TCP ACKs). Slices point into the Offload's internal buffers
|
||||||
// or we're out of rxBuf headroom. This amortizes the poll wake over
|
// and are only valid until the next Read or Close on this Queue.
|
||||||
// bursts of small packets (e.g. TCP ACKs). Packet.Bytes slices point
|
func (r *Offload) Read() ([][]byte, error) {
|
||||||
// into the Offload's internal buffer and are only valid until the next
|
|
||||||
// Read or Close on this Queue.
|
|
||||||
func (r *Offload) Read() ([]Packet, error) {
|
|
||||||
r.pending = r.pending[:0]
|
r.pending = r.pending[:0]
|
||||||
r.rxOff = 0
|
r.segOff = 0
|
||||||
|
|
||||||
// Initial (blocking) read. Retry on decode errors so a single bad
|
// Initial (blocking) read. Retry on decode errors so a single bad
|
||||||
// packet does not stall the reader.
|
// packet does not stall the reader.
|
||||||
for {
|
for {
|
||||||
n, err := r.readPacket(true)
|
n, err := r.readRaw(r.readBuf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -256,10 +196,10 @@ func (r *Offload) Read() ([]Packet, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Drain: non-blocking reads until the kernel queue is empty, the drain
|
// Drain: non-blocking reads until the kernel queue is empty, the drain
|
||||||
// cap is reached, or rxBuf no longer has room for another worst-case
|
// cap is reached, or segBuf no longer has room for another worst-case
|
||||||
// kernel-supplied packet (tunRxBufSize).
|
// superpacket.
|
||||||
for len(r.pending) < tunDrainCap && tunRxBufCap-r.rxOff >= tunRxBufSize {
|
for len(r.pending) < tunDrainCap && tunSegBufCap-r.segOff >= tunSegBufSize {
|
||||||
n, err := r.readPacket(false)
|
n, err := unix.Read(r.fd, r.readBuf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// EAGAIN / EINTR / anything else: stop draining. We already
|
// EAGAIN / EINTR / anything else: stop draining. We already
|
||||||
// have a valid batch from the first read.
|
// have a valid batch from the first read.
|
||||||
@@ -278,57 +218,21 @@ func (r *Offload) Read() ([]Packet, error) {
|
|||||||
return r.pending, nil
|
return r.pending, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// decodeRead processes the packet sitting in rxBuf at rxOff (length
|
// decodeRead decodes the virtio header plus payload in r.readBuf[:n], appends
|
||||||
// pktLen). The bytes stay in rxBuf — for GSO_NONE we slice them as a
|
// the segments to r.pending, and advances r.segOff by the total scratch used.
|
||||||
// regular IP datagram (running finishChecksum if NEEDS_CSUM is set);
|
func (r *Offload) decodeRead(n int) error {
|
||||||
// for TSO/USO superpackets we attach the corrected GSO metadata so the
|
if n < virtioNetHdrLen {
|
||||||
// caller can segment lazily at encrypt time. rxOff advances past the
|
return fmt.Errorf("short tun read: %d < %d", n, virtioNetHdrLen)
|
||||||
// kernel-supplied body and nothing else, since segmentation no longer
|
|
||||||
// writes back into rxBuf.
|
|
||||||
func (r *Offload) decodeRead(pktLen int) error {
|
|
||||||
if pktLen <= 0 {
|
|
||||||
return fmt.Errorf("short tun read: %d", pktLen)
|
|
||||||
}
|
}
|
||||||
var hdr virtio.Hdr
|
var hdr VirtioNetHdr
|
||||||
hdr.Decode(r.readVnetScratch[:])
|
hdr.decode(r.readBuf[:virtioNetHdrLen])
|
||||||
|
before := len(r.pending)
|
||||||
body := r.rxBuf[r.rxOff : r.rxOff+pktLen]
|
if err := segmentInto(r.readBuf[virtioNetHdrLen:n], hdr, &r.pending, r.segBuf[r.segOff:]); err != nil {
|
||||||
|
|
||||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
|
||||||
if err := virtio.FinishChecksum(body, hdr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
r.pending = append(r.pending, Packet{Bytes: body})
|
|
||||||
r.rxOff += pktLen
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// GSO superpacket: validate, fix the kernel-supplied HdrLen on the
|
|
||||||
// FORWARD path (CorrectHdrLen), pick the L4 protocol, and attach
|
|
||||||
// the metadata. The bytes stay in rxBuf untouched, segmentation
|
|
||||||
// happens in SegmentSuperpacket at encrypt time.
|
|
||||||
if err := virtio.CheckValid(body, hdr); err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := virtio.CorrectHdrLen(body, &hdr); err != nil {
|
for k := before; k < len(r.pending); k++ {
|
||||||
return err
|
r.segOff += len(r.pending[k])
|
||||||
}
|
}
|
||||||
proto, err := protoFromGSOType(hdr.GSOType)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
r.pending = append(r.pending, Packet{
|
|
||||||
Bytes: body,
|
|
||||||
GSO: GSOInfo{
|
|
||||||
Size: hdr.GSOSize,
|
|
||||||
HdrLen: hdr.HdrLen,
|
|
||||||
CsumStart: hdr.CsumStart,
|
|
||||||
Proto: proto,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
r.rxOff += pktLen
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -337,7 +241,7 @@ func (r *Offload) Write(buf []byte) (int, error) {
|
|||||||
{Base: &validVnetHdr[0]},
|
{Base: &validVnetHdr[0]},
|
||||||
{Base: &buf[0]},
|
{Base: &buf[0]},
|
||||||
}
|
}
|
||||||
iovs[0].SetLen(virtio.Size)
|
iovs[0].SetLen(virtioNetHdrLen)
|
||||||
iovs[1].SetLen(len(buf))
|
iovs[1].SetLen(len(buf))
|
||||||
return r.writeWithScratch(buf, &iovs)
|
return r.writeWithScratch(buf, &iovs)
|
||||||
}
|
}
|
||||||
@@ -346,6 +250,8 @@ func (r *Offload) writeWithScratch(buf []byte, iovs *[2]unix.Iovec) (int, error)
|
|||||||
if len(buf) == 0 {
|
if len(buf) == 0 {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
// Point the payload iovec at the caller's buffer. iovs[0] is pre-wired
|
||||||
|
// to validVnetHdr during Offload construction so we don't rebuild it here.
|
||||||
iovs[1].Base = &buf[0]
|
iovs[1].Base = &buf[0]
|
||||||
iovs[1].SetLen(len(buf))
|
iovs[1].SetLen(len(buf))
|
||||||
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
|
return r.rawWrite(unsafe.Slice(&iovs[0], len(iovs)))
|
||||||
@@ -355,10 +261,10 @@ func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
|||||||
for {
|
for {
|
||||||
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
|
n, _, errno := syscall.Syscall(unix.SYS_WRITEV, uintptr(r.fd), uintptr(unsafe.Pointer(&iovs[0])), uintptr(len(iovs)))
|
||||||
if errno == 0 {
|
if errno == 0 {
|
||||||
if int(n) < virtio.Size {
|
if int(n) < virtioNetHdrLen {
|
||||||
return 0, io.ErrShortWrite
|
return 0, io.ErrShortWrite
|
||||||
}
|
}
|
||||||
return int(n) - virtio.Size, nil
|
return int(n) - virtioNetHdrLen, nil
|
||||||
}
|
}
|
||||||
if errno == unix.EAGAIN {
|
if errno == unix.EAGAIN {
|
||||||
if err := r.blockOnWrite(); err != nil {
|
if err := r.blockOnWrite(); err != nil {
|
||||||
@@ -376,44 +282,29 @@ func (r *Offload) rawWrite(iovs []unix.Iovec) (int, error) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Capabilities reports the offload features negotiated for this Queue. TSO
|
// GSOSupported reports whether this queue was opened with IFF_VNET_HDR and
|
||||||
// is always true for Offload (we only construct it on IFF_VNET_HDR FDs);
|
// can accept WriteGSO. When false, callers should fall back to per-segment
|
||||||
// USO is true only when the kernel agreed to TUN_F_USO4|6 at open time
|
// Write calls.
|
||||||
// (Linux ≥ 6.2).
|
func (r *Offload) GSOSupported() bool { return true }
|
||||||
func (r *Offload) Capabilities() Capabilities {
|
|
||||||
return Capabilities{TSO: true, USO: r.usoEnabled}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error {
|
func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte) error {
|
||||||
if len(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 {
|
if len(hdr) == 0 || len(pays) == 0 || len(transportHdr) == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
// L4 checksum offset inside transportHdr: TCP=16 (the `check` field after
|
vhdr := VirtioNetHdr{
|
||||||
// seq/ack/dataoff/flags/window), UDP=6 (after sport/dport/length).
|
|
||||||
var csumOff uint16
|
|
||||||
switch proto {
|
|
||||||
case GSOProtoUDP:
|
|
||||||
csumOff = 6
|
|
||||||
default:
|
|
||||||
csumOff = 16
|
|
||||||
}
|
|
||||||
vhdr := virtio.Hdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
HdrLen: uint16(len(hdr) + len(transportHdr)),
|
HdrLen: uint16(len(hdr) + len(transportHdr)),
|
||||||
GSOSize: uint16(len(pays[0])),
|
GSOSize: uint16(len(pays[0])),
|
||||||
CsumStart: uint16(len(hdr)),
|
CsumStart: uint16(len(hdr)),
|
||||||
CsumOffset: csumOff,
|
CsumOffset: 16, // TCP checksum field lives 16 bytes into the TCP header
|
||||||
}
|
}
|
||||||
if len(pays) > 1 {
|
if len(pays) > 1 {
|
||||||
ipVer := hdr[0] >> 4
|
ipVer := hdr[0] >> 4
|
||||||
switch {
|
if ipVer == 6 {
|
||||||
case proto == GSOProtoUDP && (ipVer == 4 || ipVer == 6):
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_UDP_L4
|
|
||||||
case ipVer == 6:
|
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV6
|
||||||
case ipVer == 4:
|
} else if ipVer == 4 {
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||||
default:
|
} else {
|
||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||||
vhdr.GSOSize = 0
|
vhdr.GSOSize = 0
|
||||||
}
|
}
|
||||||
@@ -421,17 +312,18 @@ func (r *Offload) WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto
|
|||||||
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
vhdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||||
vhdr.GSOSize = 0
|
vhdr.GSOSize = 0
|
||||||
}
|
}
|
||||||
vhdr.Encode(r.gsoHdrBuf[:])
|
vhdr.encode(r.gsoHdrBuf[:])
|
||||||
|
|
||||||
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
|
// Build the iovec array: [virtio_hdr, hdr, transportHdr, pays...]. r.gsoIovs[0] is
|
||||||
// wired to gsoHdrBuf at construction and never changes.
|
// wired to gsoHdrBuf at construction and never changes.
|
||||||
need := 3 + len(pays)
|
need := 3 + len(pays)
|
||||||
if need > cap(r.gsoIovs) {
|
if cap(r.gsoIovs) < need {
|
||||||
slog.Default().Warn("tio: WriteGSO iovec budget exceeded; dropping superpacket",
|
grown := make([]unix.Iovec, need)
|
||||||
"need", need, "cap", cap(r.gsoIovs), "segments", len(pays))
|
grown[0] = r.gsoIovs[0]
|
||||||
return fmt.Errorf("tio: WriteGSO needs %d iovecs but cap is %d", need, cap(r.gsoIovs))
|
r.gsoIovs = grown
|
||||||
|
} else {
|
||||||
|
r.gsoIovs = r.gsoIovs[:need]
|
||||||
}
|
}
|
||||||
r.gsoIovs = r.gsoIovs[:need]
|
|
||||||
r.gsoIovs[1].Base = &hdr[0]
|
r.gsoIovs[1].Base = &hdr[0]
|
||||||
r.gsoIovs[1].SetLen(len(hdr))
|
r.gsoIovs[1].SetLen(len(hdr))
|
||||||
r.gsoIovs[2].Base = &transportHdr[0]
|
r.gsoIovs[2].Base = &transportHdr[0]
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ type Poll struct {
|
|||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
|
|
||||||
readBuf []byte
|
readBuf []byte
|
||||||
batchRet [1]Packet
|
batchRet [1][]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||||
@@ -97,12 +97,12 @@ func (t *Poll) blockOnWrite() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *Poll) Read() ([]Packet, error) {
|
func (t *Poll) Read() ([][]byte, error) {
|
||||||
n, err := t.readOne(t.readBuf)
|
n, err := t.readOne(t.readBuf)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
t.batchRet[0] = Packet{Bytes: t.readBuf[:n]}
|
t.batchRet[0] = t.readBuf[:n]
|
||||||
return t.batchRet[:], nil
|
return t.batchRet[:], nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
||||||
// The caller takes ownership of the read fd (pass it into a QueueSet).
|
// The caller takes ownership of the read fd (pass it to newOffload / newFriend).
|
||||||
func newReadPipe(t *testing.T) int {
|
func newReadPipe(t *testing.T) int {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
var fds [2]int
|
var fds [2]int
|
||||||
@@ -29,7 +29,7 @@ func newReadPipe(t *testing.T) int {
|
|||||||
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||||
pipe1 := newReadPipe(t)
|
pipe1 := newReadPipe(t)
|
||||||
pipe2 := newReadPipe(t)
|
pipe2 := newReadPipe(t)
|
||||||
parent, err := NewPollQueueSet()
|
parent, err := NewPollContainer()
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.NoError(t, parent.Add(pipe1))
|
require.NoError(t, parent.Add(pipe1))
|
||||||
require.NoError(t, parent.Add(pipe2))
|
require.NoError(t, parent.Add(pipe2))
|
||||||
|
|||||||
@@ -4,48 +4,328 @@
|
|||||||
package tio
|
package tio
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
|
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// protoFromGSOType maps a virtio_net_hdr GSOType to the GSOProto value the
|
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
|
||||||
// segment-time helpers use. Returns an error for GSO_NONE or any unknown
|
const (
|
||||||
// value — the caller should only invoke this on a confirmed superpacket.
|
ipv4HeaderMinLen = 20 // IHL=5, no options
|
||||||
func protoFromGSOType(t uint8) (GSOProto, error) {
|
ipv4HeaderMaxLen = 60 // IHL=15, max options
|
||||||
switch t {
|
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
|
||||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
tcpHeaderMinLen = 20 // data-offset=5, no options
|
||||||
return GSOProtoTCP, nil
|
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
||||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
)
|
||||||
return GSOProtoUDP, nil
|
|
||||||
|
// Byte offsets inside an IPv4 header.
|
||||||
|
const (
|
||||||
|
ipv4TotalLenOff = 2
|
||||||
|
ipv4IDOff = 4
|
||||||
|
ipv4ChecksumOff = 10
|
||||||
|
ipv4SrcOff = 12
|
||||||
|
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
|
||||||
|
)
|
||||||
|
|
||||||
|
// Byte offsets inside an IPv6 header.
|
||||||
|
const (
|
||||||
|
ipv6PayloadLenOff = 4
|
||||||
|
ipv6SrcOff = 8
|
||||||
|
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
|
||||||
|
)
|
||||||
|
|
||||||
|
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
|
||||||
|
const (
|
||||||
|
tcpSeqOff = 4
|
||||||
|
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
|
||||||
|
tcpFlagsOff = 13
|
||||||
|
tcpChecksumOff = 16
|
||||||
|
)
|
||||||
|
|
||||||
|
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
||||||
|
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
||||||
|
|
||||||
|
func checkVirtioValid(pkt []byte, hdr VirtioNetHdr) error {
|
||||||
|
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
|
||||||
|
// carry coalescing info rather than checksum offsets. A TUN writing via
|
||||||
|
// IFF_VNET_HDR should never emit this, but if it did we would silently
|
||||||
|
// miscompute the segment checksums — refuse the packet instead.
|
||||||
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
||||||
|
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
||||||
|
}
|
||||||
|
if len(pkt) < ipv4HeaderMinLen {
|
||||||
|
return fmt.Errorf("packet too short")
|
||||||
|
}
|
||||||
|
ipVersion := pkt[0] >> 4
|
||||||
|
switch hdr.GSOType {
|
||||||
|
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
||||||
|
if ipVersion != 4 {
|
||||||
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||||
|
}
|
||||||
|
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||||
|
if ipVersion != 6 {
|
||||||
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||||
|
}
|
||||||
default:
|
default:
|
||||||
return 0, fmt.Errorf("unsupported virtio gso type: %d", t)
|
if !(ipVersion == 6 || ipVersion == 4) {
|
||||||
|
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleGSONone(pkt []byte, hdr VirtioNetHdr, out *[][]byte, scratch []byte) error {
|
||||||
|
if len(pkt) > len(scratch) {
|
||||||
|
return fmt.Errorf("packet larger than segment buffer: %d > %d", len(pkt), len(scratch))
|
||||||
|
}
|
||||||
|
copy(scratch, pkt)
|
||||||
|
seg := scratch[:len(pkt)]
|
||||||
|
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
||||||
|
if err := finishChecksum(seg, hdr); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
*out = append(*out, seg)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func correctHdrLen(pkt []byte, hdr *VirtioNetHdr) error {
|
||||||
|
// Thank you wireguard-go for documenting these edge-cases
|
||||||
|
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
|
||||||
|
// of the entire first packet when the kernel is handling it as part of a
|
||||||
|
// FORWARD path. Instead, parse the transport header length and add it onto
|
||||||
|
// csumStart, which is synonymous for IP header length.
|
||||||
|
const tcpDataOffset = 12
|
||||||
|
|
||||||
|
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
||||||
|
hdr.HdrLen = hdr.CsumStart + 8
|
||||||
|
} else {
|
||||||
|
if len(pkt) <= int(hdr.CsumStart+tcpDataOffset) {
|
||||||
|
return errors.New("packet is too short")
|
||||||
|
}
|
||||||
|
|
||||||
|
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffset] >> 4 * 4)
|
||||||
|
if tcpHLen < 20 || tcpHLen > 60 {
|
||||||
|
// A TCP header must be between 20 and 60 bytes in length.
|
||||||
|
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
||||||
|
}
|
||||||
|
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(pkt) < int(hdr.HdrLen) {
|
||||||
|
return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen)
|
||||||
|
}
|
||||||
|
|
||||||
|
if hdr.HdrLen < hdr.CsumStart {
|
||||||
|
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
|
||||||
|
}
|
||||||
|
cSumAt := int(hdr.CsumStart + hdr.CsumStart)
|
||||||
|
if cSumAt+1 >= len(pkt) {
|
||||||
|
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// segmentInto splits a TUN-side packet described by hdr into one or more
|
||||||
|
// IP packets, each appended to *out as a slice of scratch. scratch must be
|
||||||
|
// sized to hold every segment (including replicated headers).
|
||||||
|
func segmentInto(pkt []byte, hdr VirtioNetHdr, out *[][]byte, scratch []byte) error {
|
||||||
|
if err := checkVirtioValid(pkt, hdr); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
||||||
|
return handleGSONone(pkt, hdr, out, scratch)
|
||||||
|
}
|
||||||
|
if err := correctHdrLen(pkt, &hdr); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
switch hdr.GSOType {
|
||||||
|
case unix.VIRTIO_NET_HDR_GSO_TCPV4, unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
||||||
|
return segmentTCP(pkt, hdr, out, scratch)
|
||||||
|
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported virtio gso type: %d", hdr.GSOType)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// SegmentSuperpacket invokes fn once per segment of pkt. For non-GSO pkts
|
// finishChecksum computes the L4 checksum for a non-GSO packet that the kernel
|
||||||
// fn is called once with pkt.Bytes (no segmentation, no copy). For GSO/USO
|
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit
|
||||||
// superpackets fn is called once per segment with a slice of pkt.Bytes
|
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
|
||||||
// holding that segment's plaintext (a freshly-patched L3+L4 header sliced
|
// the pseudo-header partial sum by the kernel), and store the result.
|
||||||
// in front of the original payload chunk). The slide is destructive: pkt is
|
func finishChecksum(seg []byte, hdr VirtioNetHdr) error {
|
||||||
// consumed by this call and its bytes are in an undefined state when
|
cs := int(hdr.CsumStart)
|
||||||
// SegmentSuperpacket returns. Callers must not retain pkt or any earlier
|
co := int(hdr.CsumOffset)
|
||||||
// seg slice past fn's return for that segment. The scratch parameter is
|
if cs+co+2 > len(seg) {
|
||||||
// unused on the destructive path and kept only for cross-platform
|
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
|
||||||
// signature compatibility. Aborts and returns the first error from fn or
|
|
||||||
// from per-segment construction.
|
|
||||||
func SegmentSuperpacket(pkt Packet, fn func(seg []byte) error) error {
|
|
||||||
if !pkt.GSO.IsSuperpacket() {
|
|
||||||
return fn(pkt.Bytes)
|
|
||||||
}
|
|
||||||
switch pkt.GSO.Proto {
|
|
||||||
case GSOProtoTCP:
|
|
||||||
return virtio.SegmentTCP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
|
||||||
case GSOProtoUDP:
|
|
||||||
return virtio.SegmentUDP(pkt.Bytes, pkt.GSO.HdrLen, pkt.GSO.CsumStart, pkt.GSO.Size, fn)
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("unsupported gso proto: %d", pkt.GSO.Proto)
|
|
||||||
}
|
}
|
||||||
|
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
|
||||||
|
// L4 region starting at cs, folding the prior partial in as the seed.
|
||||||
|
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
||||||
|
seg[cs+co] = 0
|
||||||
|
seg[cs+co+1] = 0
|
||||||
|
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// segmentTCP software-segments a TSO superpacket into one IP packet per MSS
|
||||||
|
// chunk. The caller guarantees hdr.GSOType is TCPV4 or TCPV6.
|
||||||
|
//
|
||||||
|
// Hot-path shape: the per-segment loop only sums the payload chunk. The TCP
|
||||||
|
// header, the IPv4 header, and the pseudo-header src/dst/proto contributions
|
||||||
|
// are each summed once up front — every segment reuses those three pre-folded
|
||||||
|
// uint32 values and combines them with small per-segment deltas (seq, flags,
|
||||||
|
// tcpLen, ip_id, total_len) that are cheap to fold in.
|
||||||
|
func segmentTCP(pkt []byte, hdr VirtioNetHdr, out *[][]byte, scratch []byte) error {
|
||||||
|
if hdr.GSOSize == 0 {
|
||||||
|
return fmt.Errorf("gso_size is zero")
|
||||||
|
}
|
||||||
|
if hdr.CsumStart == 0 {
|
||||||
|
return fmt.Errorf("csum_start is zero")
|
||||||
|
}
|
||||||
|
|
||||||
|
isV4 := hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_TCPV4
|
||||||
|
headerLen := int(hdr.HdrLen) // already corrected by the caller
|
||||||
|
csumStart := int(hdr.CsumStart)
|
||||||
|
|
||||||
|
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
||||||
|
|
||||||
|
payload := pkt[headerLen:]
|
||||||
|
payLen := len(payload)
|
||||||
|
gsoSize := int(hdr.GSOSize)
|
||||||
|
numSeg := (payLen + gsoSize - 1) / gsoSize
|
||||||
|
if numSeg == 0 {
|
||||||
|
numSeg = 1
|
||||||
|
}
|
||||||
|
|
||||||
|
need := numSeg*headerLen + payLen
|
||||||
|
if need > len(scratch) {
|
||||||
|
return fmt.Errorf("scratch too small for %d segments: need %d have %d", numSeg, need, len(scratch))
|
||||||
|
}
|
||||||
|
|
||||||
|
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
||||||
|
origFlags := pkt[csumStart+tcpFlagsOff]
|
||||||
|
|
||||||
|
// Precompute the TCP header sum with seq/flags/csum zeroed. Copy onto
|
||||||
|
// the stack, zero the per-segment-varying fields, sum once.
|
||||||
|
var tmp [tcpHeaderMaxLen]byte
|
||||||
|
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen])
|
||||||
|
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
||||||
|
tmp[tcpFlagsOff] = 0
|
||||||
|
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
||||||
|
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
|
||||||
|
|
||||||
|
// Pseudo-header src+dst+proto contribution (tcpLen varies per segment).
|
||||||
|
var baseProtoSum uint32
|
||||||
|
if isV4 {
|
||||||
|
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
||||||
|
} else {
|
||||||
|
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
||||||
|
}
|
||||||
|
baseProtoSum += uint32(unix.IPPROTO_TCP)
|
||||||
|
|
||||||
|
// Precompute IPv4 header sum with total_len/id/csum zeroed.
|
||||||
|
var origIPID uint16
|
||||||
|
var ihl int
|
||||||
|
var baseIPHdrSum uint32
|
||||||
|
if isV4 {
|
||||||
|
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
||||||
|
ihl = int(pkt[0]&0x0f) * 4
|
||||||
|
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
||||||
|
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
||||||
|
}
|
||||||
|
var ipTmp [ipv4HeaderMaxLen]byte
|
||||||
|
copy(ipTmp[:ihl], pkt[:ihl])
|
||||||
|
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
||||||
|
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
||||||
|
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
off := 0
|
||||||
|
for i := 0; i < numSeg; i++ {
|
||||||
|
segStart := i * gsoSize
|
||||||
|
segEnd := segStart + gsoSize
|
||||||
|
if segEnd > payLen {
|
||||||
|
segEnd = payLen
|
||||||
|
}
|
||||||
|
segPayLen := segEnd - segStart
|
||||||
|
|
||||||
|
copy(scratch[off:], pkt[:headerLen])
|
||||||
|
copy(scratch[off+headerLen:], payload[segStart:segEnd])
|
||||||
|
seg := scratch[off : off+headerLen+segPayLen]
|
||||||
|
off += headerLen + segPayLen
|
||||||
|
|
||||||
|
segSeq := origSeq + uint32(segStart)
|
||||||
|
segFlags := origFlags
|
||||||
|
if i != numSeg-1 {
|
||||||
|
segFlags = origFlags &^ tcpFinPshMask
|
||||||
|
}
|
||||||
|
totalLen := headerLen + segPayLen
|
||||||
|
|
||||||
|
// Patch IP header and write the v4 header checksum from the precomputed base.
|
||||||
|
if isV4 {
|
||||||
|
segID := origIPID + uint16(i)
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
||||||
|
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
||||||
|
} else {
|
||||||
|
// IPv6 payload length excludes the fixed header but includes any
|
||||||
|
// extension headers between [ipv6FixedLen:csumStart].
|
||||||
|
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Patch TCP header.
|
||||||
|
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
|
||||||
|
seg[csumStart+tcpFlagsOff] = segFlags
|
||||||
|
// (csum is written below; its prior contents in `seg` don't affect the
|
||||||
|
// computation since we never sum over the segment's own header.)
|
||||||
|
|
||||||
|
tcpLen := tcpHdrLen + segPayLen
|
||||||
|
paySum := uint32(checksum.Checksum(payload[segStart:segEnd], 0))
|
||||||
|
|
||||||
|
// Combine pre-folded uint32s into a wider accumulator, then fold. Using
|
||||||
|
// uint64 guards against overflow when segSeq's high bits set.
|
||||||
|
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
||||||
|
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
||||||
|
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||||
|
wide = (wide & 0xffffffff) + (wide >> 32)
|
||||||
|
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
|
||||||
|
|
||||||
|
*out = append(*out, seg)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
||||||
|
// complements it, yielding the on-wire Internet checksum value.
|
||||||
|
func foldComplement(sum uint32) uint16 {
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
sum = (sum & 0xffff) + (sum >> 16)
|
||||||
|
return ^uint16(sum)
|
||||||
|
}
|
||||||
|
|
||||||
|
// pseudoHeaderIPv4 returns the folded pseudo-header sum used to verify a TCP
|
||||||
|
// segment's checksum in tests. src/dst are 4 bytes each.
|
||||||
|
func pseudoHeaderIPv4(src, dst []byte, proto byte, tcpLen int) uint16 {
|
||||||
|
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
||||||
|
s += uint32(proto) + uint32(tcpLen)
|
||||||
|
s = (s & 0xffff) + (s >> 16)
|
||||||
|
s = (s & 0xffff) + (s >> 16)
|
||||||
|
return uint16(s)
|
||||||
|
}
|
||||||
|
|
||||||
|
// pseudoHeaderIPv6 returns the folded pseudo-header sum used to verify a TCP
|
||||||
|
// segment's checksum in tests. src/dst are 16 bytes each.
|
||||||
|
func pseudoHeaderIPv6(src, dst []byte, proto byte, tcpLen int) uint16 {
|
||||||
|
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
||||||
|
s += uint32(tcpLen>>16) + uint32(tcpLen&0xffff) + uint32(proto)
|
||||||
|
s = (s & 0xffff) + (s >> 16)
|
||||||
|
s = (s & 0xffff) + (s >> 16)
|
||||||
|
return uint16(s)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,78 +10,17 @@ import (
|
|||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/tio/virtio"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// testSegScratchSize is a generous segmentation scratch sized to fit any
|
|
||||||
// of the synthetic TSO/USO superpackets these tests generate (one
|
|
||||||
// worst-case 64 KiB superpacket plus replicated per-segment headers).
|
|
||||||
const testSegScratchSize = 192 * 1024
|
|
||||||
|
|
||||||
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
// verifyChecksum confirms that the one's-complement sum across `b`, seeded
|
||||||
// with a folded pseudo-header sum, equals all-ones (valid).
|
// with a folded pseudo-header sum, equals all-ones (valid).
|
||||||
func verifyChecksum(b []byte, pseudo uint16) bool {
|
func verifyChecksum(b []byte, pseudo uint16) bool {
|
||||||
return checksum.Checksum(b, pseudo) == 0xffff
|
return checksum.Checksum(b, pseudo) == 0xffff
|
||||||
}
|
}
|
||||||
|
|
||||||
// segmentForTest is the test-only counterpart to the production
|
|
||||||
// SegmentSuperpacket path. It handles GSO_NONE (with optional
|
|
||||||
// finishChecksum) inline and dispatches GSO superpackets through
|
|
||||||
// SegmentSuperpacket, draining each yielded segment into a
|
|
||||||
// freshly-copied [][]byte slot so callers can iterate after the call
|
|
||||||
// returns. Tests pre-set hdr.HdrLen correctly, so correctHdrLen is not
|
|
||||||
// invoked here.
|
|
||||||
func segmentForTest(pkt []byte, hdr virtio.Hdr, out *[][]byte, scratch []byte) error {
|
|
||||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_NONE {
|
|
||||||
cp := append([]byte(nil), pkt...)
|
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_NEEDS_CSUM != 0 {
|
|
||||||
if err := virtio.FinishChecksum(cp, hdr); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
*out = append(*out, cp)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
proto, err := protoFromGSOType(hdr.GSOType)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
gso := GSOInfo{
|
|
||||||
Size: hdr.GSOSize,
|
|
||||||
HdrLen: hdr.HdrLen,
|
|
||||||
CsumStart: hdr.CsumStart,
|
|
||||||
Proto: proto,
|
|
||||||
}
|
|
||||||
return SegmentSuperpacket(Packet{Bytes: pkt, GSO: gso}, func(seg []byte) error {
|
|
||||||
*out = append(*out, append([]byte(nil), seg...))
|
|
||||||
return nil
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// pseudoHeaderIPv4 returns the folded pseudo-header sum used to verify a
|
|
||||||
// TCP/UDP segment's checksum in tests. src/dst are 4 bytes each.
|
|
||||||
func pseudoHeaderIPv4(src, dst []byte, proto byte, l4Len int) uint16 {
|
|
||||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
|
||||||
s += uint32(proto) + uint32(l4Len)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
return uint16(s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pseudoHeaderIPv6 returns the folded pseudo-header sum used to verify a
|
|
||||||
// TCP/UDP segment's checksum in tests. src/dst are 16 bytes each.
|
|
||||||
func pseudoHeaderIPv6(src, dst []byte, proto byte, l4Len int) uint16 {
|
|
||||||
s := uint32(checksum.Checksum(src, 0)) + uint32(checksum.Checksum(dst, 0))
|
|
||||||
s += uint32(l4Len>>16) + uint32(l4Len&0xffff) + uint32(proto)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
s = (s & 0xffff) + (s >> 16)
|
|
||||||
return uint16(s)
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildTSOv4 builds a synthetic IPv4/TCP TSO superpacket with a payload of
|
// buildTSOv4 builds a synthetic IPv4/TCP TSO superpacket with a payload of
|
||||||
// `payLen` bytes split at `mss`.
|
// `payLen` bytes split at `mss`.
|
||||||
func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
|
func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, VirtioNetHdr) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
const ipLen = 20
|
const ipLen = 20
|
||||||
const tcpLen = 20
|
const tcpLen = 20
|
||||||
@@ -111,7 +50,7 @@ func buildTSOv4(t *testing.T, payLen, mss int) ([]byte, virtio.Hdr) {
|
|||||||
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
pkt[ipLen+tcpLen+i] = byte(i & 0xff)
|
||||||
}
|
}
|
||||||
|
|
||||||
return pkt, virtio.Hdr{
|
return pkt, VirtioNetHdr{
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
HdrLen: uint16(ipLen + tcpLen),
|
||||||
@@ -126,10 +65,10 @@ func TestSegmentTCPv4(t *testing.T) {
|
|||||||
const numSeg = 3
|
const numSeg = 3
|
||||||
pkt, hdr := buildTSOv4(t, mss*numSeg, mss)
|
pkt, hdr := buildTSOv4(t, mss*numSeg, mss)
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, tunSegBufSize)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
t.Fatalf("segmentTCP: %v", err)
|
||||||
}
|
}
|
||||||
if len(out) != numSeg {
|
if len(out) != numSeg {
|
||||||
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
||||||
@@ -175,10 +114,10 @@ func TestSegmentTCPv4(t *testing.T) {
|
|||||||
func TestSegmentTCPv4OddTail(t *testing.T) {
|
func TestSegmentTCPv4OddTail(t *testing.T) {
|
||||||
// Payload of 250 bytes with MSS 100 → segments of 100, 100, 50.
|
// Payload of 250 bytes with MSS 100 → segments of 100, 100, 50.
|
||||||
pkt, hdr := buildTSOv4(t, 250, 100)
|
pkt, hdr := buildTSOv4(t, 250, 100)
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, tunSegBufSize)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
t.Fatalf("segmentTCP: %v", err)
|
||||||
}
|
}
|
||||||
if len(out) != 3 {
|
if len(out) != 3 {
|
||||||
t.Fatalf("want 3 segments, got %d", len(out))
|
t.Fatalf("want 3 segments, got %d", len(out))
|
||||||
@@ -232,7 +171,7 @@ func TestSegmentTCPv6(t *testing.T) {
|
|||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
}
|
}
|
||||||
|
|
||||||
hdr := virtio.Hdr{
|
hdr := VirtioNetHdr{
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
HdrLen: uint16(ipLen + tcpLen),
|
||||||
@@ -241,10 +180,10 @@ func TestSegmentTCPv6(t *testing.T) {
|
|||||||
CsumOffset: 16,
|
CsumOffset: 16,
|
||||||
}
|
}
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, tunSegBufSize)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
t.Fatalf("segmentTCP: %v", err)
|
||||||
}
|
}
|
||||||
if len(out) != numSeg {
|
if len(out) != numSeg {
|
||||||
t.Fatalf("want %d segments, got %d", numSeg, len(out))
|
t.Fatalf("want %d segments, got %d", numSeg, len(out))
|
||||||
@@ -284,10 +223,10 @@ func TestSegmentGSONonePassesThrough(t *testing.T) {
|
|||||||
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
hdr.GSOType = unix.VIRTIO_NET_HDR_GSO_NONE
|
||||||
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
hdr.Flags = 0 // no NEEDS_CSUM, leave packet untouched
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, tunSegBufSize)
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
if err := segmentInto(pkt, hdr, &out, scratch); err != nil {
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
t.Fatalf("segmentInto: %v", err)
|
||||||
}
|
}
|
||||||
if len(out) != 1 {
|
if len(out) != 1 {
|
||||||
t.Fatalf("want 1 segment, got %d", len(out))
|
t.Fatalf("want 1 segment, got %d", len(out))
|
||||||
@@ -297,254 +236,11 @@ func TestSegmentGSONonePassesThrough(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestSegmentRejectsLegacyUDPGSO ensures the legacy GSO_UDP (UFO) marker is
|
func TestSegmentRejectsUDP(t *testing.T) {
|
||||||
// still rejected; only modern GSO_UDP_L4 (USO) is supported.
|
hdr := VirtioNetHdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
|
||||||
func TestSegmentRejectsLegacyUDPGSO(t *testing.T) {
|
|
||||||
hdr := virtio.Hdr{GSOType: unix.VIRTIO_NET_HDR_GSO_UDP}
|
|
||||||
var out [][]byte
|
var out [][]byte
|
||||||
if err := segmentForTest(nil, hdr, &out, nil); err == nil {
|
if err := segmentInto(nil, hdr, &out, nil); err == nil {
|
||||||
t.Fatalf("expected rejection for legacy UDP GSO")
|
t.Fatalf("expected rejection for UDP GSO")
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// buildUSOv4 builds a synthetic IPv4/UDP USO superpacket with payload of
|
|
||||||
// payLen bytes, segmented at gsoSize.
|
|
||||||
func buildUSOv4(t *testing.T, payLen, gsoSize int) ([]byte, virtio.Hdr) {
|
|
||||||
t.Helper()
|
|
||||||
const ipLen = 20
|
|
||||||
const udpLen = 8
|
|
||||||
pkt := make([]byte, ipLen+udpLen+payLen)
|
|
||||||
|
|
||||||
// IPv4 header
|
|
||||||
pkt[0] = 0x45 // version 4, IHL 5
|
|
||||||
binary.BigEndian.PutUint16(pkt[2:4], uint16(ipLen+udpLen+payLen))
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], 0x4242)
|
|
||||||
pkt[8] = 64
|
|
||||||
pkt[9] = unix.IPPROTO_UDP
|
|
||||||
copy(pkt[12:16], []byte{10, 0, 0, 1})
|
|
||||||
copy(pkt[16:20], []byte{10, 0, 0, 2})
|
|
||||||
|
|
||||||
// UDP header (length + checksum filled in per segment by segmentUDPYield)
|
|
||||||
binary.BigEndian.PutUint16(pkt[20:22], 12345) // sport
|
|
||||||
binary.BigEndian.PutUint16(pkt[22:24], 53) // dport
|
|
||||||
|
|
||||||
for i := 0; i < payLen; i++ {
|
|
||||||
pkt[ipLen+udpLen+i] = byte(i & 0xff)
|
|
||||||
}
|
|
||||||
|
|
||||||
return pkt, virtio.Hdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
|
||||||
HdrLen: uint16(ipLen + udpLen),
|
|
||||||
GSOSize: uint16(gsoSize),
|
|
||||||
CsumStart: uint16(ipLen),
|
|
||||||
CsumOffset: 6,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentUDPv4(t *testing.T) {
|
|
||||||
const gso = 100
|
|
||||||
const numSeg = 3
|
|
||||||
pkt, hdr := buildUSOv4(t, gso*numSeg, gso)
|
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != numSeg {
|
|
||||||
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, seg := range out {
|
|
||||||
if len(seg) != 28+gso {
|
|
||||||
t.Errorf("seg %d: len %d want %d", i, len(seg), 28+gso)
|
|
||||||
}
|
|
||||||
totalLen := binary.BigEndian.Uint16(seg[2:4])
|
|
||||||
if totalLen != uint16(28+gso) {
|
|
||||||
t.Errorf("seg %d: total_len=%d want %d", i, totalLen, 28+gso)
|
|
||||||
}
|
|
||||||
// kernel UDP-GSO does NOT bump the IPv4 ID across segments; every
|
|
||||||
// segment carries the same ID as the seed.
|
|
||||||
id := binary.BigEndian.Uint16(seg[4:6])
|
|
||||||
if id != 0x4242 {
|
|
||||||
t.Errorf("seg %d: ip id=%#x want %#x", i, id, 0x4242)
|
|
||||||
}
|
|
||||||
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
|
||||||
if udpLen != uint16(8+gso) {
|
|
||||||
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+gso)
|
|
||||||
}
|
|
||||||
if !verifyChecksum(seg[:20], 0) {
|
|
||||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, 8+gso)
|
|
||||||
if !verifyChecksum(seg[20:], psum) {
|
|
||||||
t.Errorf("seg %d: bad UDP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentUDPv4OddTail(t *testing.T) {
|
|
||||||
// 250 bytes payload, gsoSize=100 → segments of 100, 100, 50.
|
|
||||||
pkt, hdr := buildUSOv4(t, 250, 100)
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != 3 {
|
|
||||||
t.Fatalf("want 3 segments, got %d", len(out))
|
|
||||||
}
|
|
||||||
wantPay := []int{100, 100, 50}
|
|
||||||
for i, seg := range out {
|
|
||||||
if len(seg)-28 != wantPay[i] {
|
|
||||||
t.Errorf("seg %d: pay len %d want %d", i, len(seg)-28, wantPay[i])
|
|
||||||
}
|
|
||||||
udpLen := binary.BigEndian.Uint16(seg[24:26])
|
|
||||||
if udpLen != uint16(8+wantPay[i]) {
|
|
||||||
t.Errorf("seg %d: udp len=%d want %d", i, udpLen, 8+wantPay[i])
|
|
||||||
}
|
|
||||||
if !verifyChecksum(seg[:20], 0) {
|
|
||||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_UDP, 8+wantPay[i])
|
|
||||||
if !verifyChecksum(seg[20:], psum) {
|
|
||||||
t.Errorf("seg %d: bad UDP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSegmentUDPv6(t *testing.T) {
|
|
||||||
const ipLen = 40
|
|
||||||
const udpLen = 8
|
|
||||||
const gso = 120
|
|
||||||
const numSeg = 2
|
|
||||||
payLen := gso * numSeg
|
|
||||||
pkt := make([]byte, ipLen+udpLen+payLen)
|
|
||||||
|
|
||||||
// IPv6 header
|
|
||||||
pkt[0] = 0x60
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(udpLen+payLen))
|
|
||||||
pkt[6] = unix.IPPROTO_UDP
|
|
||||||
pkt[7] = 64
|
|
||||||
pkt[8] = 0xfe
|
|
||||||
pkt[9] = 0x80
|
|
||||||
pkt[23] = 1
|
|
||||||
pkt[24] = 0xfe
|
|
||||||
pkt[25] = 0x80
|
|
||||||
pkt[39] = 2
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 53)
|
|
||||||
|
|
||||||
for i := 0; i < payLen; i++ {
|
|
||||||
pkt[ipLen+udpLen+i] = byte(i)
|
|
||||||
}
|
|
||||||
|
|
||||||
hdr := virtio.Hdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_UDP_L4,
|
|
||||||
HdrLen: uint16(ipLen + udpLen),
|
|
||||||
GSOSize: uint16(gso),
|
|
||||||
CsumStart: uint16(ipLen),
|
|
||||||
CsumOffset: 6,
|
|
||||||
}
|
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != numSeg {
|
|
||||||
t.Fatalf("want %d segments, got %d", numSeg, len(out))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, seg := range out {
|
|
||||||
if len(seg) != ipLen+udpLen+gso {
|
|
||||||
t.Errorf("seg %d: len %d want %d", i, len(seg), ipLen+udpLen+gso)
|
|
||||||
}
|
|
||||||
pl := binary.BigEndian.Uint16(seg[4:6])
|
|
||||||
if pl != uint16(udpLen+gso) {
|
|
||||||
t.Errorf("seg %d: payload_length=%d want %d", i, pl, udpLen+gso)
|
|
||||||
}
|
|
||||||
ul := binary.BigEndian.Uint16(seg[ipLen+4 : ipLen+6])
|
|
||||||
if ul != uint16(udpLen+gso) {
|
|
||||||
t.Errorf("seg %d: udp len=%d want %d", i, ul, udpLen+gso)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_UDP, udpLen+gso)
|
|
||||||
if !verifyChecksum(seg[ipLen:], psum) {
|
|
||||||
t.Errorf("seg %d: bad UDP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSegmentUDPCEPropagates confirms IP-level CE marks on the seed appear on
|
|
||||||
// every segment. UDP has no transport-level CWR/ECE: the IP TOS/TC byte is
|
|
||||||
// copied verbatim into every segment by the segment-prefix copy.
|
|
||||||
func TestSegmentUDPCEPropagates(t *testing.T) {
|
|
||||||
pkt, hdr := buildUSOv4(t, 200, 100)
|
|
||||||
pkt[1] = 0x03 // CE codepoint in IP-ECN
|
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != 2 {
|
|
||||||
t.Fatalf("want 2 segments, got %d", len(out))
|
|
||||||
}
|
|
||||||
for i, seg := range out {
|
|
||||||
if seg[1]&0x03 != 0x03 {
|
|
||||||
t.Errorf("seg %d: CE missing (tos=%#x)", i, seg[1])
|
|
||||||
}
|
|
||||||
if !verifyChecksum(seg[:20], 0) {
|
|
||||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestSegmentTCPCwrFirstSegmentOnly confirms RFC 3168 §6.1.2: when a TSO
|
|
||||||
// burst's seed has CWR set, only the first emitted segment carries CWR.
|
|
||||||
// ECE is preserved on every segment (different signal, persistent state).
|
|
||||||
func TestSegmentTCPCwrFirstSegmentOnly(t *testing.T) {
|
|
||||||
const mss = 100
|
|
||||||
const numSeg = 3
|
|
||||||
pkt, hdr := buildTSOv4(t, mss*numSeg, mss)
|
|
||||||
// Seed flags: CWR | ECE | ACK | PSH.
|
|
||||||
pkt[33] = 0x80 | 0x40 | 0x10 | 0x08
|
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
|
||||||
var out [][]byte
|
|
||||||
if err := segmentForTest(pkt, hdr, &out, scratch); err != nil {
|
|
||||||
t.Fatalf("segmentForTest: %v", err)
|
|
||||||
}
|
|
||||||
if len(out) != numSeg {
|
|
||||||
t.Fatalf("expected %d segments, got %d", numSeg, len(out))
|
|
||||||
}
|
|
||||||
for i, seg := range out {
|
|
||||||
flags := seg[33]
|
|
||||||
hasCwr := flags&0x80 != 0
|
|
||||||
hasEce := flags&0x40 != 0
|
|
||||||
hasPsh := flags&0x08 != 0
|
|
||||||
wantCwr := i == 0
|
|
||||||
wantPsh := i == numSeg-1
|
|
||||||
if hasCwr != wantCwr {
|
|
||||||
t.Errorf("seg %d: CWR=%v want %v (flags=%#x)", i, hasCwr, wantCwr, flags)
|
|
||||||
}
|
|
||||||
if !hasEce {
|
|
||||||
t.Errorf("seg %d: ECE missing (flags=%#x)", i, flags)
|
|
||||||
}
|
|
||||||
if hasPsh != wantPsh {
|
|
||||||
t.Errorf("seg %d: PSH=%v want %v (flags=%#x)", i, hasPsh, wantPsh, flags)
|
|
||||||
}
|
|
||||||
// IP and TCP checksums must still verify after the flag rewrite.
|
|
||||||
if !verifyChecksum(seg[:20], 0) {
|
|
||||||
t.Errorf("seg %d: bad IPv4 header checksum", i)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv4(seg[12:16], seg[16:20], unix.IPPROTO_TCP, 20+mss)
|
|
||||||
if !verifyChecksum(seg[20:], psum) {
|
|
||||||
t.Errorf("seg %d: bad TCP checksum", i)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -580,7 +276,7 @@ func BenchmarkSegmentTCPv4(b *testing.B) {
|
|||||||
for i := 0; i < sz.payLen; i++ {
|
for i := 0; i < sz.payLen; i++ {
|
||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
pkt[ipLen+tcpLen+i] = byte(i)
|
||||||
}
|
}
|
||||||
hdr := virtio.Hdr{
|
hdr := VirtioNetHdr{
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV4,
|
||||||
HdrLen: uint16(ipLen + tcpLen),
|
HdrLen: uint16(ipLen + tcpLen),
|
||||||
@@ -589,23 +285,14 @@ func BenchmarkSegmentTCPv4(b *testing.B) {
|
|||||||
CsumOffset: 16,
|
CsumOffset: 16,
|
||||||
}
|
}
|
||||||
|
|
||||||
scratch := make([]byte, testSegScratchSize)
|
scratch := make([]byte, tunSegBufSize)
|
||||||
out := make([][]byte, 0, 64)
|
out := make([][]byte, 0, 64)
|
||||||
|
|
||||||
// SegmentSuperpacket consumes its input destructively; restore
|
|
||||||
// pkt from a master copy each iteration. The restore mirrors the
|
|
||||||
// kernel→userspace copy that hands a fresh GSO blob to the
|
|
||||||
// segmenter in production, so it's representative cost rather
|
|
||||||
// than bench overhead.
|
|
||||||
master := append([]byte(nil), pkt...)
|
|
||||||
work := make([]byte, len(pkt))
|
|
||||||
|
|
||||||
b.SetBytes(int64(len(pkt)))
|
b.SetBytes(int64(len(pkt)))
|
||||||
b.ResetTimer()
|
b.ResetTimer()
|
||||||
for i := 0; i < b.N; i++ {
|
for i := 0; i < b.N; i++ {
|
||||||
copy(work, master)
|
|
||||||
out = out[:0]
|
out = out[:0]
|
||||||
if err := segmentForTest(work, hdr, &out, scratch); err != nil {
|
if err := segmentTCP(pkt, hdr, &out, scratch); err != nil {
|
||||||
b.Fatal(err)
|
b.Fatal(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -639,156 +326,3 @@ func TestTunFileWriteVnetHdrNoAlloc(t *testing.T) {
|
|||||||
t.Fatalf("Write allocated %.1f times per call, want 0", allocs)
|
t.Fatalf("Write allocated %.1f times per call, want 0", allocs)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildTSOv6 builds a synthetic IPv6/TCP TSO superpacket with payLen bytes
|
|
||||||
// of payload, segmented at gso. Returns the packet bytes only; the
|
|
||||||
// virtio_net_hdr is the caller's responsibility.
|
|
||||||
func buildTSOv6(payLen, gso int) []byte {
|
|
||||||
const ipLen = 40
|
|
||||||
const tcpLen = 20
|
|
||||||
pkt := make([]byte, ipLen+tcpLen+payLen)
|
|
||||||
|
|
||||||
pkt[0] = 0x60 // version 6
|
|
||||||
binary.BigEndian.PutUint16(pkt[4:6], uint16(tcpLen+payLen))
|
|
||||||
pkt[6] = unix.IPPROTO_TCP
|
|
||||||
pkt[7] = 64
|
|
||||||
pkt[8] = 0xfe
|
|
||||||
pkt[9] = 0x80
|
|
||||||
pkt[23] = 1
|
|
||||||
pkt[24] = 0xfe
|
|
||||||
pkt[25] = 0x80
|
|
||||||
pkt[39] = 2
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(pkt[40:42], 12345)
|
|
||||||
binary.BigEndian.PutUint16(pkt[42:44], 80)
|
|
||||||
binary.BigEndian.PutUint32(pkt[44:48], 7)
|
|
||||||
binary.BigEndian.PutUint32(pkt[48:52], 99)
|
|
||||||
pkt[52] = 0x50
|
|
||||||
pkt[53] = 0x10 // ACK only
|
|
||||||
binary.BigEndian.PutUint16(pkt[54:56], 65535)
|
|
||||||
|
|
||||||
for i := 0; i < payLen; i++ {
|
|
||||||
pkt[ipLen+tcpLen+i] = byte(i)
|
|
||||||
}
|
|
||||||
return pkt
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestDecodeReadFitsMaxTSOAtDrainThreshold proves the rxBuf sizing is
|
|
||||||
// correct: when rxOff is at the maximum value the drain headroom check
|
|
||||||
// allows, decodeRead must still be able to absorb a worst-case 64KiB
|
|
||||||
// TSO superpacket without dropping the burst. With segmentation deferred
|
|
||||||
// to encrypt time, decodeRead writes only the kernel-supplied bytes into
|
|
||||||
// rxBuf, so the size requirement is just "fit one worst-case input."
|
|
||||||
//
|
|
||||||
// Regression history: in a prior layout the rx buffer doubled as the
|
|
||||||
// segmentation output, a near-threshold drain read returned "scratch too
|
|
||||||
// small", the whole 45-segment TSO burst was dropped, and the remote's TCP
|
|
||||||
// fast-retransmit collapsed cwnd. Keeping this test in the new layout
|
|
||||||
// guards against re-introducing a drain headroom shortfall.
|
|
||||||
func TestDecodeReadFitsMaxTSOAtDrainThreshold(t *testing.T) {
|
|
||||||
const ipv6HdrLen = 40
|
|
||||||
const tcpHdrLen = 20
|
|
||||||
const headerLen = ipv6HdrLen + tcpHdrLen
|
|
||||||
// Maximum TUN read body. The tunReadBufSize cap on readv's body iovec
|
|
||||||
// is what bounds the kernel's superpacket length.
|
|
||||||
pktLen := tunReadBufSize
|
|
||||||
payLen := pktLen - headerLen
|
|
||||||
const targetSegs = 64
|
|
||||||
gsoSize := (payLen + targetSegs - 1) / targetSegs
|
|
||||||
|
|
||||||
pkt := buildTSOv6(payLen, gsoSize)
|
|
||||||
if len(pkt) != pktLen {
|
|
||||||
t.Fatalf("buildTSOv6 produced %d bytes, want %d", len(pkt), pktLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
o := &Offload{
|
|
||||||
rxBuf: make([]byte, tunRxBufCap),
|
|
||||||
}
|
|
||||||
// rxOff at the maximum value the drain headroom check permits before
|
|
||||||
// it would refuse another read. Any drain-time read up to this
|
|
||||||
// threshold MUST still process correctly.
|
|
||||||
o.rxOff = tunRxBufCap - tunRxBufSize
|
|
||||||
|
|
||||||
// Stage the body in rxBuf as if readv(2) just placed it there.
|
|
||||||
copy(o.rxBuf[o.rxOff:], pkt)
|
|
||||||
|
|
||||||
// Encode the matching virtio_net_hdr.
|
|
||||||
hdr := virtio.Hdr{
|
|
||||||
Flags: unix.VIRTIO_NET_HDR_F_NEEDS_CSUM,
|
|
||||||
GSOType: unix.VIRTIO_NET_HDR_GSO_TCPV6,
|
|
||||||
HdrLen: uint16(headerLen),
|
|
||||||
GSOSize: uint16(gsoSize),
|
|
||||||
CsumStart: uint16(ipv6HdrLen),
|
|
||||||
CsumOffset: 16,
|
|
||||||
}
|
|
||||||
hdr.Encode(o.readVnetScratch[:])
|
|
||||||
|
|
||||||
startRxOff := o.rxOff
|
|
||||||
if err := o.decodeRead(pktLen); err != nil {
|
|
||||||
t.Fatalf("decodeRead at drain threshold returned %v — rxBuf sizing regression: "+
|
|
||||||
"tunRxBufSize=%d must hold one worst-case input (%d)",
|
|
||||||
err, tunRxBufSize, pktLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(o.pending) != 1 {
|
|
||||||
t.Fatalf("got %d packets, want 1 superpacket entry", len(o.pending))
|
|
||||||
}
|
|
||||||
got := o.pending[0]
|
|
||||||
if !got.GSO.IsSuperpacket() {
|
|
||||||
t.Fatalf("expected superpacket GSO metadata, got %+v", got.GSO)
|
|
||||||
}
|
|
||||||
if got.GSO.Proto != GSOProtoTCP {
|
|
||||||
t.Errorf("GSO.Proto=%d want TCP", got.GSO.Proto)
|
|
||||||
}
|
|
||||||
if got.GSO.Size != uint16(gsoSize) {
|
|
||||||
t.Errorf("GSO.Size=%d want %d", got.GSO.Size, gsoSize)
|
|
||||||
}
|
|
||||||
if got.GSO.HdrLen != uint16(headerLen) {
|
|
||||||
t.Errorf("GSO.HdrLen=%d want %d", got.GSO.HdrLen, headerLen)
|
|
||||||
}
|
|
||||||
if got.GSO.CsumStart != uint16(ipv6HdrLen) {
|
|
||||||
t.Errorf("GSO.CsumStart=%d want %d", got.GSO.CsumStart, ipv6HdrLen)
|
|
||||||
}
|
|
||||||
if len(got.Bytes) != pktLen {
|
|
||||||
t.Errorf("len(Bytes)=%d want %d", len(got.Bytes), pktLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
// rxOff advances exactly by the kernel-supplied body length — no
|
|
||||||
// segmentation output to account for any more.
|
|
||||||
if o.rxOff != startRxOff+pktLen {
|
|
||||||
t.Errorf("rxOff=%d want %d", o.rxOff, startRxOff+pktLen)
|
|
||||||
}
|
|
||||||
if o.rxOff > tunRxBufCap {
|
|
||||||
t.Fatalf("rxOff=%d overran rxBuf (cap=%d)", o.rxOff, tunRxBufCap)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Validate that segmenting the returned superpacket reproduces the
|
|
||||||
// expected per-segment IPv6 payload length and TCP checksum.
|
|
||||||
wantSegs := (payLen + gsoSize - 1) / gsoSize
|
|
||||||
gotSegs := 0
|
|
||||||
if err := SegmentSuperpacket(got, func(seg []byte) error {
|
|
||||||
defer func() { gotSegs++ }()
|
|
||||||
if len(seg) < headerLen+1 {
|
|
||||||
t.Errorf("seg %d too short: %d", gotSegs, len(seg))
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
if seg[0]>>4 != 6 {
|
|
||||||
t.Errorf("seg %d: bad IP version %#x", gotSegs, seg[0])
|
|
||||||
}
|
|
||||||
segPay := len(seg) - headerLen
|
|
||||||
gotPL := binary.BigEndian.Uint16(seg[4:6])
|
|
||||||
if gotPL != uint16(tcpHdrLen+segPay) {
|
|
||||||
t.Errorf("seg %d: payload_len=%d want %d", gotSegs, gotPL, tcpHdrLen+segPay)
|
|
||||||
}
|
|
||||||
psum := pseudoHeaderIPv6(seg[8:24], seg[24:40], unix.IPPROTO_TCP, tcpHdrLen+segPay)
|
|
||||||
if !verifyChecksum(seg[ipv6HdrLen:], psum) {
|
|
||||||
t.Errorf("seg %d: bad TCP checksum", gotSegs)
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}); err != nil {
|
|
||||||
t.Fatalf("SegmentSuperpacket: %v", err)
|
|
||||||
}
|
|
||||||
if gotSegs != wantSegs {
|
|
||||||
t.Fatalf("got %d segments, want %d", gotSegs, wantSegs)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,402 +0,0 @@
|
|||||||
//go:build linux && !android
|
|
||||||
// +build linux,!android
|
|
||||||
|
|
||||||
// Package virtio implements the pure validation, header-correction, and
|
|
||||||
// per-segment slicing logic for kernel-supplied TSO/USO superpackets on
|
|
||||||
// IFF_VNET_HDR TUN devices. It is FD-free and depends only on the byte
|
|
||||||
// layout of the virtio_net_hdr and the IP/TCP/UDP headers it describes,
|
|
||||||
// so it can be unit-tested in isolation from the tio Queue runtime.
|
|
||||||
package virtio
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
|
|
||||||
"github.com/slackhq/nebula/overlay/checksum"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Protocol header size bounds used to validate / cap kernel-supplied offsets.
|
|
||||||
const (
|
|
||||||
ipv4HeaderMinLen = 20 // IHL=5, no options
|
|
||||||
ipv4HeaderMaxLen = 60 // IHL=15, max options
|
|
||||||
ipv6FixedLen = 40 // IPv6 base header; extensions would extend this
|
|
||||||
tcpHeaderMinLen = 20 // data-offset=5, no options
|
|
||||||
tcpHeaderMaxLen = 60 // data-offset=15, max options
|
|
||||||
)
|
|
||||||
|
|
||||||
// Byte offsets inside an IPv4 header.
|
|
||||||
const (
|
|
||||||
ipv4TotalLenOff = 2
|
|
||||||
ipv4IDOff = 4
|
|
||||||
ipv4ChecksumOff = 10
|
|
||||||
ipv4SrcOff = 12
|
|
||||||
ipv4AddrsEnd = 20 // end of dst address (ipv4SrcOff + 2*4)
|
|
||||||
)
|
|
||||||
|
|
||||||
// Byte offsets inside an IPv6 header.
|
|
||||||
const (
|
|
||||||
ipv6PayloadLenOff = 4
|
|
||||||
ipv6SrcOff = 8
|
|
||||||
ipv6AddrsEnd = 40 // end of dst address (ipv6SrcOff + 2*16)
|
|
||||||
)
|
|
||||||
|
|
||||||
// Byte offsets inside a TCP header (relative to its start, i.e. csumStart).
|
|
||||||
const (
|
|
||||||
tcpSeqOff = 4
|
|
||||||
tcpDataOffOff = 12 // upper nibble is header len in 32-bit words
|
|
||||||
tcpFlagsOff = 13
|
|
||||||
tcpChecksumOff = 16
|
|
||||||
)
|
|
||||||
|
|
||||||
// UDP header is fixed at 8 bytes: {sport, dport, length, checksum}.
|
|
||||||
const (
|
|
||||||
udpHeaderLen = 8
|
|
||||||
udpLengthOff = 4
|
|
||||||
udpChecksumOff = 6
|
|
||||||
)
|
|
||||||
|
|
||||||
// tcpFinPshMask is cleared on every segment except the last of a TSO burst.
|
|
||||||
const tcpFinPshMask = 0x09 // FIN(0x01) | PSH(0x08)
|
|
||||||
|
|
||||||
// tcpCwrFlag is cleared on every segment except the first. Per RFC 3168
|
|
||||||
// §6.1.2 the CWR bit signals a one-shot transition (the sender just halved
|
|
||||||
// its window) and must appear on the first segment of a TSO burst only.
|
|
||||||
const tcpCwrFlag = 0x80
|
|
||||||
|
|
||||||
// CheckValid rejects packets whose virtio_net_hdr/IP combination would
|
|
||||||
// cause a downstream miscompute. The TUN should never emit RSC_INFO and
|
|
||||||
// the GSO type must agree with the IP version nibble.
|
|
||||||
func CheckValid(pkt []byte, hdr Hdr) error {
|
|
||||||
// When RSC_INFO is set the csum_start/csum_offset fields are repurposed to
|
|
||||||
// carry coalescing info rather than checksum offsets. A TUN writing via
|
|
||||||
// IFF_VNET_HDR should never emit this, but if it did we would silently
|
|
||||||
// miscompute the segment checksums — refuse the packet instead.
|
|
||||||
if hdr.Flags&unix.VIRTIO_NET_HDR_F_RSC_INFO != 0 {
|
|
||||||
return fmt.Errorf("virtio RSC_INFO flag not supported on TUN reads")
|
|
||||||
}
|
|
||||||
if len(pkt) < ipv4HeaderMinLen {
|
|
||||||
return fmt.Errorf("packet too short")
|
|
||||||
}
|
|
||||||
ipVersion := pkt[0] >> 4
|
|
||||||
switch hdr.GSOType {
|
|
||||||
case unix.VIRTIO_NET_HDR_GSO_TCPV4:
|
|
||||||
if ipVersion != 4 {
|
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
|
||||||
}
|
|
||||||
case unix.VIRTIO_NET_HDR_GSO_TCPV6:
|
|
||||||
if ipVersion != 6 {
|
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
|
||||||
}
|
|
||||||
case unix.VIRTIO_NET_HDR_GSO_UDP_L4:
|
|
||||||
// USO carries either v4 or v6; the leading nibble disambiguates.
|
|
||||||
if !(ipVersion == 4 || ipVersion == 6) {
|
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
if !(ipVersion == 6 || ipVersion == 4) {
|
|
||||||
return fmt.Errorf("invalid IP version %d for GSO type %d", ipVersion, hdr.GSOType)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// CorrectHdrLen rewrites hdr.HdrLen based on the actual transport header
|
|
||||||
// length read out of pkt. The kernel's hdr.HdrLen on the FORWARD path can
|
|
||||||
// be the length of the entire first packet, so we don't trust it.
|
|
||||||
func CorrectHdrLen(pkt []byte, hdr *Hdr) error {
|
|
||||||
// Thank you wireguard-go for documenting these edge-cases
|
|
||||||
// Don't trust hdr.hdrLen from the kernel as it can be equal to the length
|
|
||||||
// of the entire first packet when the kernel is handling it as part of a
|
|
||||||
// FORWARD path. Instead, parse the transport header length and add it onto
|
|
||||||
// csumStart, which is synonymous for IP header length.
|
|
||||||
|
|
||||||
if hdr.GSOType == unix.VIRTIO_NET_HDR_GSO_UDP_L4 {
|
|
||||||
hdr.HdrLen = hdr.CsumStart + 8
|
|
||||||
} else {
|
|
||||||
if len(pkt) <= int(hdr.CsumStart+tcpDataOffOff) {
|
|
||||||
return errors.New("packet is too short")
|
|
||||||
}
|
|
||||||
|
|
||||||
tcpHLen := uint16(pkt[hdr.CsumStart+tcpDataOffOff] >> 4 * 4)
|
|
||||||
if tcpHLen < 20 || tcpHLen > 60 {
|
|
||||||
// A TCP header must be between 20 and 60 bytes in length.
|
|
||||||
return fmt.Errorf("tcp header len is invalid: %d", tcpHLen)
|
|
||||||
}
|
|
||||||
hdr.HdrLen = hdr.CsumStart + tcpHLen
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(pkt) < int(hdr.HdrLen) {
|
|
||||||
return fmt.Errorf("length of packet (%d) < virtioNetHdr.HdrLen (%d)", len(pkt), hdr.HdrLen)
|
|
||||||
}
|
|
||||||
|
|
||||||
if hdr.HdrLen < hdr.CsumStart {
|
|
||||||
return fmt.Errorf("virtioNetHdr.HdrLen (%d) < virtioNetHdr.CsumStart (%d)", hdr.HdrLen, hdr.CsumStart)
|
|
||||||
}
|
|
||||||
cSumAt := int(hdr.CsumStart + hdr.CsumStart)
|
|
||||||
if cSumAt+1 >= len(pkt) {
|
|
||||||
return fmt.Errorf("end of checksum offset (%d) exceeds packet length (%d)", cSumAt+1, len(pkt))
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// SegmentTCP walks a TSO superpacket pkt, yielding each segment as a
|
|
||||||
// slice into pkt itself. Per-segment plaintext is laid out by sliding a
|
|
||||||
// freshly-patched copy of the L3+L4 header into pkt at offset i*gsoSize,
|
|
||||||
// where it sits immediately before that segment's payload chunk in the
|
|
||||||
// original buffer. The slide is destructive: iter i's header write overwrites
|
|
||||||
// the last hdrLen bytes of seg_{i-1}'s payload, which is dead by the time
|
|
||||||
// the next iteration begins. pkt is consumed by this call and must not be
|
|
||||||
// inspected by the caller after the final yield.
|
|
||||||
func SegmentTCP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
|
||||||
if gsoSizeU == 0 {
|
|
||||||
return fmt.Errorf("gso_size is zero")
|
|
||||||
}
|
|
||||||
if csumStartU == 0 {
|
|
||||||
return fmt.Errorf("csum_start is zero")
|
|
||||||
}
|
|
||||||
|
|
||||||
headerLen := int(hdrLenU)
|
|
||||||
csumStart := int(csumStartU)
|
|
||||||
isV4 := pkt[0]>>4 == 4
|
|
||||||
|
|
||||||
tcpHdrLen := int(pkt[csumStart+tcpDataOffOff]>>4) * 4
|
|
||||||
payLen := len(pkt) - headerLen
|
|
||||||
gsoSize := int(gsoSizeU)
|
|
||||||
numSeg := (payLen + gsoSize - 1) / gsoSize
|
|
||||||
if numSeg == 0 {
|
|
||||||
numSeg = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
origSeq := binary.BigEndian.Uint32(pkt[csumStart+tcpSeqOff : csumStart+tcpSeqOff+4])
|
|
||||||
origFlags := pkt[csumStart+tcpFlagsOff]
|
|
||||||
|
|
||||||
var tmp [tcpHeaderMaxLen]byte
|
|
||||||
copy(tmp[:tcpHdrLen], pkt[csumStart:headerLen])
|
|
||||||
tmp[tcpSeqOff], tmp[tcpSeqOff+1], tmp[tcpSeqOff+2], tmp[tcpSeqOff+3] = 0, 0, 0, 0
|
|
||||||
tmp[tcpFlagsOff] = 0
|
|
||||||
tmp[tcpChecksumOff], tmp[tcpChecksumOff+1] = 0, 0
|
|
||||||
baseTcpHdrSum := uint32(checksum.Checksum(tmp[:tcpHdrLen], 0))
|
|
||||||
|
|
||||||
var baseProtoSum uint32
|
|
||||||
if isV4 {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
|
||||||
} else {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
|
||||||
}
|
|
||||||
baseProtoSum += uint32(unix.IPPROTO_TCP)
|
|
||||||
|
|
||||||
var origIPID uint16
|
|
||||||
var baseIPHdrSum uint32
|
|
||||||
if isV4 {
|
|
||||||
origIPID = binary.BigEndian.Uint16(pkt[ipv4IDOff : ipv4IDOff+2])
|
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
|
||||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
|
||||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
|
||||||
}
|
|
||||||
var ipTmp [ipv4HeaderMaxLen]byte
|
|
||||||
copy(ipTmp[:ihl], pkt[:ihl])
|
|
||||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4IDOff], ipTmp[ipv4IDOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
|
||||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := 0; i < numSeg; i++ {
|
|
||||||
segStart := i * gsoSize
|
|
||||||
segEnd := segStart + gsoSize
|
|
||||||
if segEnd > payLen {
|
|
||||||
segEnd = payLen
|
|
||||||
}
|
|
||||||
segPayLen := segEnd - segStart
|
|
||||||
segLen := headerLen + segPayLen
|
|
||||||
headerOff := i * gsoSize
|
|
||||||
|
|
||||||
// Slide the header into place immediately before this segment's
|
|
||||||
// payload. Iter 0's header is already at pkt[:headerLen]; for
|
|
||||||
// i ≥ 1 we copy from there. The constant-byte fields of pkt[:headerLen]
|
|
||||||
// survive iter 0's in-place patches (only seq/flags/cksum/totalLen/id
|
|
||||||
// are touched), and iter 0's stale variable-field values are
|
|
||||||
// overwritten by the per-segment patches below.
|
|
||||||
if i > 0 {
|
|
||||||
copy(pkt[headerOff:headerOff+headerLen], pkt[:headerLen])
|
|
||||||
}
|
|
||||||
seg := pkt[headerOff : headerOff+segLen]
|
|
||||||
|
|
||||||
segSeq := origSeq + uint32(segStart)
|
|
||||||
segFlags := origFlags
|
|
||||||
if i != 0 {
|
|
||||||
segFlags &^= tcpCwrFlag
|
|
||||||
}
|
|
||||||
if i != numSeg-1 {
|
|
||||||
segFlags &^= tcpFinPshMask
|
|
||||||
}
|
|
||||||
totalLen := segLen
|
|
||||||
|
|
||||||
if isV4 {
|
|
||||||
segID := origIPID + uint16(i)
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4IDOff:ipv4IDOff+2], segID)
|
|
||||||
ipSum := baseIPHdrSum + uint32(totalLen) + uint32(segID)
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
|
||||||
} else {
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
|
||||||
}
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint32(seg[csumStart+tcpSeqOff:csumStart+tcpSeqOff+4], segSeq)
|
|
||||||
seg[csumStart+tcpFlagsOff] = segFlags
|
|
||||||
|
|
||||||
tcpLen := tcpHdrLen + segPayLen
|
|
||||||
// Payload bytes still live at their original offset in pkt. The
|
|
||||||
// header slide above only writes into pkt[i*G : i*G+H], which is
|
|
||||||
// the tail of seg_{i-1}'s payload (already consumed) and never
|
|
||||||
// overlaps seg_i's own payload at pkt[H+i*G : H+(i+1)*G].
|
|
||||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
|
||||||
wide := uint64(baseTcpHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
|
||||||
wide += uint64(segSeq) + uint64(segFlags) + uint64(tcpLen)
|
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
|
||||||
binary.BigEndian.PutUint16(seg[csumStart+tcpChecksumOff:csumStart+tcpChecksumOff+2], foldComplement(uint32(wide)))
|
|
||||||
|
|
||||||
if err := yield(seg); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// SegmentUDP walks a USO superpacket, sliding a per-segment-patched
|
|
||||||
// L3+L4 header into pkt at offset i*gsoSize and yielding pkt[i*G:i*G+segLen]
|
|
||||||
// to the caller. Per-segment patches are total_len + IPv4 csum (or IPv6
|
|
||||||
// payload_len) plus the UDP length and checksum. pkt is consumed
|
|
||||||
// destructively; see SegmentTCP for the layout reasoning.
|
|
||||||
//
|
|
||||||
// UDP-GSO leaves the IPv4 ID identical across segments (the kernel does not
|
|
||||||
// bump it), which is why the IP-level per-segment work is limited to
|
|
||||||
// total_len + IPv4 header checksum (v4) or payload_len (v6).
|
|
||||||
func SegmentUDP(pkt []byte, hdrLenU, csumStartU, gsoSizeU uint16, yield func(seg []byte) error) error {
|
|
||||||
if gsoSizeU == 0 {
|
|
||||||
return fmt.Errorf("gso_size is zero")
|
|
||||||
}
|
|
||||||
if csumStartU == 0 {
|
|
||||||
return fmt.Errorf("csum_start is zero")
|
|
||||||
}
|
|
||||||
|
|
||||||
isV4 := pkt[0]>>4 == 4
|
|
||||||
headerLen := int(hdrLenU)
|
|
||||||
csumStart := int(csumStartU)
|
|
||||||
if headerLen-csumStart != udpHeaderLen {
|
|
||||||
return fmt.Errorf("udp header len mismatch: %d", headerLen-csumStart)
|
|
||||||
}
|
|
||||||
|
|
||||||
payLen := len(pkt) - headerLen
|
|
||||||
gsoSize := int(gsoSizeU)
|
|
||||||
numSeg := (payLen + gsoSize - 1) / gsoSize
|
|
||||||
if numSeg == 0 {
|
|
||||||
numSeg = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
var udpTmp [udpHeaderLen]byte
|
|
||||||
copy(udpTmp[:], pkt[csumStart:headerLen])
|
|
||||||
udpTmp[udpLengthOff], udpTmp[udpLengthOff+1] = 0, 0
|
|
||||||
udpTmp[udpChecksumOff], udpTmp[udpChecksumOff+1] = 0, 0
|
|
||||||
baseUDPHdrSum := uint32(checksum.Checksum(udpTmp[:], 0))
|
|
||||||
|
|
||||||
var baseProtoSum uint32
|
|
||||||
if isV4 {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv4SrcOff:ipv4AddrsEnd], 0))
|
|
||||||
} else {
|
|
||||||
baseProtoSum = uint32(checksum.Checksum(pkt[ipv6SrcOff:ipv6AddrsEnd], 0))
|
|
||||||
}
|
|
||||||
baseProtoSum += uint32(unix.IPPROTO_UDP)
|
|
||||||
|
|
||||||
var baseIPHdrSum uint32
|
|
||||||
if isV4 {
|
|
||||||
ihl := int(pkt[0]&0x0f) * 4
|
|
||||||
if ihl < ipv4HeaderMinLen || ihl > csumStart {
|
|
||||||
return fmt.Errorf("bad IPv4 IHL: %d", ihl)
|
|
||||||
}
|
|
||||||
var ipTmp [ipv4HeaderMaxLen]byte
|
|
||||||
copy(ipTmp[:ihl], pkt[:ihl])
|
|
||||||
ipTmp[ipv4TotalLenOff], ipTmp[ipv4TotalLenOff+1] = 0, 0
|
|
||||||
ipTmp[ipv4ChecksumOff], ipTmp[ipv4ChecksumOff+1] = 0, 0
|
|
||||||
baseIPHdrSum = uint32(checksum.Checksum(ipTmp[:ihl], 0))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := 0; i < numSeg; i++ {
|
|
||||||
segStart := i * gsoSize
|
|
||||||
segEnd := segStart + gsoSize
|
|
||||||
if segEnd > payLen {
|
|
||||||
segEnd = payLen
|
|
||||||
}
|
|
||||||
segPayLen := segEnd - segStart
|
|
||||||
segLen := headerLen + segPayLen
|
|
||||||
headerOff := i * gsoSize
|
|
||||||
|
|
||||||
if i > 0 {
|
|
||||||
copy(pkt[headerOff:headerOff+headerLen], pkt[:headerLen])
|
|
||||||
}
|
|
||||||
seg := pkt[headerOff : headerOff+segLen]
|
|
||||||
|
|
||||||
totalLen := segLen
|
|
||||||
udpLen := udpHeaderLen + segPayLen
|
|
||||||
|
|
||||||
if isV4 {
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4TotalLenOff:ipv4TotalLenOff+2], uint16(totalLen))
|
|
||||||
ipSum := baseIPHdrSum + uint32(totalLen)
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv4ChecksumOff:ipv4ChecksumOff+2], foldComplement(ipSum))
|
|
||||||
} else {
|
|
||||||
binary.BigEndian.PutUint16(seg[ipv6PayloadLenOff:ipv6PayloadLenOff+2], uint16(headerLen-ipv6FixedLen+segPayLen))
|
|
||||||
}
|
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(seg[csumStart+udpLengthOff:csumStart+udpLengthOff+2], uint16(udpLen))
|
|
||||||
|
|
||||||
paySum := uint32(checksum.Checksum(pkt[headerLen+segStart:headerLen+segEnd], 0))
|
|
||||||
wide := uint64(baseUDPHdrSum) + uint64(paySum) + uint64(baseProtoSum)
|
|
||||||
wide += uint64(udpLen) + uint64(udpLen)
|
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
|
||||||
wide = (wide & 0xffffffff) + (wide >> 32)
|
|
||||||
csum := foldComplement(uint32(wide))
|
|
||||||
if csum == 0 {
|
|
||||||
csum = 0xffff
|
|
||||||
}
|
|
||||||
binary.BigEndian.PutUint16(seg[csumStart+udpChecksumOff:csumStart+udpChecksumOff+2], csum)
|
|
||||||
|
|
||||||
if err := yield(seg); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// FinishChecksum computes the L4 checksum for a non-GSO packet that the kernel
|
|
||||||
// handed us with NEEDS_CSUM set. csum_start / csum_offset point at the 16-bit
|
|
||||||
// checksum field; we zero it, fold a full sum (the field was pre-loaded with
|
|
||||||
// the pseudo-header partial sum by the kernel), and store the result.
|
|
||||||
func FinishChecksum(seg []byte, hdr Hdr) error {
|
|
||||||
cs := int(hdr.CsumStart)
|
|
||||||
co := int(hdr.CsumOffset)
|
|
||||||
if cs+co+2 > len(seg) {
|
|
||||||
return fmt.Errorf("csum offsets out of range: start=%d offset=%d len=%d", cs, co, len(seg))
|
|
||||||
}
|
|
||||||
// The kernel stores a partial pseudo-header sum at [cs+co:]; sum over the
|
|
||||||
// L4 region starting at cs, folding the prior partial in as the seed.
|
|
||||||
partial := binary.BigEndian.Uint16(seg[cs+co : cs+co+2])
|
|
||||||
seg[cs+co] = 0
|
|
||||||
seg[cs+co+1] = 0
|
|
||||||
binary.BigEndian.PutUint16(seg[cs+co:cs+co+2], ^checksum.Checksum(seg[cs:], partial))
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// foldComplement folds a 32-bit one's-complement partial sum to 16 bits and
|
|
||||||
// complements it, yielding the on-wire Internet checksum value.
|
|
||||||
func foldComplement(sum uint32) uint16 {
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
sum = (sum & 0xffff) + (sum >> 16)
|
|
||||||
return ^uint16(sum)
|
|
||||||
}
|
|
||||||
@@ -1,17 +1,12 @@
|
|||||||
//go:build linux && !android
|
package tio
|
||||||
// +build linux,!android
|
|
||||||
|
|
||||||
package virtio
|
|
||||||
|
|
||||||
import "encoding/binary"
|
import "encoding/binary"
|
||||||
|
|
||||||
// Size is the on-wire length of struct virtio_net_hdr the kernel
|
// Size of the legacy struct virtio_net_hdr that the kernel prepends/expects on
|
||||||
// prepends/expects on a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ
|
// a TUN opened with IFF_VNET_HDR (TUNSETVNETHDRSZ not set).
|
||||||
// not set).
|
const virtioNetHdrLen = 10
|
||||||
const Size = 10
|
|
||||||
|
|
||||||
// Hdr is the Go view of the legacy virtio_net_hdr.
|
type VirtioNetHdr struct {
|
||||||
type Hdr struct {
|
|
||||||
Flags uint8
|
Flags uint8
|
||||||
GSOType uint8
|
GSOType uint8
|
||||||
HdrLen uint16
|
HdrLen uint16
|
||||||
@@ -20,9 +15,9 @@ type Hdr struct {
|
|||||||
CsumOffset uint16
|
CsumOffset uint16
|
||||||
}
|
}
|
||||||
|
|
||||||
// Decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
// decode reads a virtio_net_hdr in host byte order (TUN default; we never
|
||||||
// call TUNSETVNETLE so the kernel matches our endianness).
|
// call TUNSETVNETLE so the kernel matches our endianness).
|
||||||
func (h *Hdr) Decode(b []byte) {
|
func (h *VirtioNetHdr) decode(b []byte) {
|
||||||
h.Flags = b[0]
|
h.Flags = b[0]
|
||||||
h.GSOType = b[1]
|
h.GSOType = b[1]
|
||||||
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
h.HdrLen = binary.NativeEndian.Uint16(b[2:4])
|
||||||
@@ -31,9 +26,10 @@ func (h *Hdr) Decode(b []byte) {
|
|||||||
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
h.CsumOffset = binary.NativeEndian.Uint16(b[8:10])
|
||||||
}
|
}
|
||||||
|
|
||||||
// Encode is the inverse of Decode: writes the virtio_net_hdr fields into b
|
// encode is the inverse of decode: writes the virtio_net_hdr fields into b
|
||||||
// (must be at least Size bytes). Used to emit a TSO superpacket on egress.
|
// (must be at least virtioNetHdrLen bytes). Used to emit a TSO superpacket
|
||||||
func (h *Hdr) Encode(b []byte) {
|
// on egress.
|
||||||
|
func (h *VirtioNetHdr) encode(b []byte) {
|
||||||
b[0] = h.Flags
|
b[0] = h.Flags
|
||||||
b[1] = h.GSOType
|
b[1] = h.GSOType
|
||||||
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
|
binary.NativeEndian.PutUint16(b[2:4], h.HdrLen)
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user