mirror of
https://github.com/slackhq/nebula.git
synced 2026-08-15 22:57:03 +02:00
Compare commits
31 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 60e07370de | |||
| 44530cb610 | |||
| c5abf30102 | |||
| 3a329ec217 | |||
| afcdf2163b | |||
| 9a30c5b6a1 | |||
| ad0c99f262 | |||
| c1ad7a3af2 | |||
| cef69465db | |||
| daef13e53a | |||
| 6e1df81473 | |||
| c61de54ec3 | |||
| 1b59636028 | |||
| 487bae4c2f | |||
| 2bdd284993 | |||
| 398d67e2da | |||
| 696903d6d9 | |||
| c82db210ef | |||
| 1ada3d4dd9 | |||
| 5f920fdd7d | |||
| cba9ea5b1f | |||
| 83809a599a | |||
| 23c67bd8d8 | |||
| dd3a7ad03c | |||
| dd2ac5d655 | |||
| 76e82a5256 | |||
| eaf756ea6c | |||
| a82a8dc547 | |||
| 213dd46588 | |||
| 4fb5cdb4fa | |||
| ff91c37529 |
@@ -0,0 +1,113 @@
|
|||||||
|
name: Code-sign Windows binaries
|
||||||
|
description: >
|
||||||
|
Sign every .exe under a given path in place via the DefinedNet code-signer
|
||||||
|
Lambda. If `role` or `bucket` is empty, logs a notice and skips signing so
|
||||||
|
forks and dev branches without AWS access still produce usable builds.
|
||||||
|
|
||||||
|
inputs:
|
||||||
|
path:
|
||||||
|
description: "Directory whose .exe files should be signed in place"
|
||||||
|
required: true
|
||||||
|
role:
|
||||||
|
description: "IAM role ARN to assume via OIDC; empty disables signing"
|
||||||
|
required: false
|
||||||
|
default: ""
|
||||||
|
bucket:
|
||||||
|
description: "S3 staging bucket the code-signer Lambda reads from; empty disables signing"
|
||||||
|
required: false
|
||||||
|
default: ""
|
||||||
|
region:
|
||||||
|
description: "AWS region for the role and Lambda"
|
||||||
|
required: false
|
||||||
|
default: "us-east-2"
|
||||||
|
function-name:
|
||||||
|
description: "Code-signer Lambda function name"
|
||||||
|
required: false
|
||||||
|
default: "code-signer"
|
||||||
|
key-prefix:
|
||||||
|
description: "S3 key prefix the caller is authorized to write under"
|
||||||
|
required: false
|
||||||
|
default: "code-signing/slackhq/nebula"
|
||||||
|
|
||||||
|
runs:
|
||||||
|
using: composite
|
||||||
|
steps:
|
||||||
|
- name: Skip notice
|
||||||
|
if: inputs.role == '' || inputs.bucket == ''
|
||||||
|
shell: sh
|
||||||
|
run: echo "::notice::code-signer role or bucket not set; skipping code signing."
|
||||||
|
|
||||||
|
- name: Configure AWS credentials
|
||||||
|
if: inputs.role != '' && inputs.bucket != ''
|
||||||
|
uses: aws-actions/configure-aws-credentials@v6
|
||||||
|
with:
|
||||||
|
role-to-assume: ${{ inputs.role }}
|
||||||
|
aws-region: ${{ inputs.region }}
|
||||||
|
# Default is 12 retries to ride out IAM trust-policy propagation; once
|
||||||
|
# the role is stable we want a real misconfiguration to fail fast.
|
||||||
|
retry-max-attempts: 5
|
||||||
|
|
||||||
|
- name: Sign .exe files
|
||||||
|
if: inputs.role != '' && inputs.bucket != ''
|
||||||
|
shell: sh
|
||||||
|
env:
|
||||||
|
SIGN_PATH: ${{ inputs.path }}
|
||||||
|
BUCKET: ${{ inputs.bucket }}
|
||||||
|
FUNCTION_NAME: ${{ inputs.function-name }}
|
||||||
|
KEY_PREFIX: ${{ inputs.key-prefix }}
|
||||||
|
run: |
|
||||||
|
set -eu
|
||||||
|
RUN="${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}"
|
||||||
|
|
||||||
|
find "$SIGN_PATH" -name '*.exe' -print | while read -r path
|
||||||
|
do
|
||||||
|
rel=${path#"$SIGN_PATH"/}
|
||||||
|
file=$(basename "$path")
|
||||||
|
name=${file%.exe}
|
||||||
|
prefix="${KEY_PREFIX}/${RUN}"
|
||||||
|
src="${prefix}/unsigned/${rel}"
|
||||||
|
dst="${prefix}/signed/${rel}"
|
||||||
|
|
||||||
|
echo "::group::Sign ${rel}"
|
||||||
|
echo "Uploading unsigned to s3://${BUCKET}/${src}"
|
||||||
|
aws s3 cp --no-progress "$path" "s3://${BUCKET}/${src}" >/dev/null
|
||||||
|
|
||||||
|
echo "Invoking ${FUNCTION_NAME} Lambda"
|
||||||
|
payload=$(jq -nc \
|
||||||
|
--arg s "$src" \
|
||||||
|
--arg d "$dst" \
|
||||||
|
--arg p "$name" \
|
||||||
|
'{source_key: $s, dest_key: $d, program_name: $p}')
|
||||||
|
meta=$(aws lambda invoke \
|
||||||
|
--function-name "$FUNCTION_NAME" \
|
||||||
|
--cli-binary-format raw-in-base64-out \
|
||||||
|
--payload "$payload" \
|
||||||
|
--output json \
|
||||||
|
/tmp/sign-resp.json)
|
||||||
|
if echo "$meta" | jq -e '.FunctionError != null' >/dev/null
|
||||||
|
then
|
||||||
|
echo "::endgroup::"
|
||||||
|
echo "::error::code-signer Lambda failed for ${rel}"
|
||||||
|
cat /tmp/sign-resp.json >&2
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
echo "Downloading signed back to ${path}"
|
||||||
|
aws s3 cp --no-progress "s3://${BUCKET}/${dst}" "$path" >/dev/null
|
||||||
|
|
||||||
|
aws s3 rm "s3://${BUCKET}/${src}" >/dev/null 2>&1 || true
|
||||||
|
aws s3 rm "s3://${BUCKET}/${dst}" >/dev/null 2>&1 || true
|
||||||
|
|
||||||
|
# Sanity-check the bytes we got back actually carry an Authenticode
|
||||||
|
# signature that this machine can validate end to end.
|
||||||
|
status=$(powershell -NoProfile -Command "(Get-AuthenticodeSignature -FilePath '$path').Status" | tr -d '\r')
|
||||||
|
if [ "$status" != "Valid" ]
|
||||||
|
then
|
||||||
|
echo "::endgroup::"
|
||||||
|
echo "::error::${rel} signature status: ${status} (expected Valid)"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
echo "Signed ${rel} (sha256=$(jq -r '.sha256' /tmp/sign-resp.json), status=${status})"
|
||||||
|
echo "::endgroup::"
|
||||||
|
done
|
||||||
@@ -24,7 +24,7 @@ jobs:
|
|||||||
mv build/*.tar.gz release
|
mv build/*.tar.gz release
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v6
|
uses: actions/upload-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: linux-latest
|
name: linux-latest
|
||||||
path: release
|
path: release
|
||||||
@@ -32,6 +32,9 @@ 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
|
||||||
|
|
||||||
@@ -54,8 +57,15 @@ 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@v6
|
uses: actions/upload-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: windows-latest
|
name: windows-latest
|
||||||
path: build
|
path: build
|
||||||
@@ -75,7 +85,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@v6
|
uses: Apple-Actions/import-codesign-certs@v7
|
||||||
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 }}
|
||||||
@@ -104,7 +114,7 @@ jobs:
|
|||||||
fi
|
fi
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v6
|
uses: actions/upload-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: darwin-latest
|
name: darwin-latest
|
||||||
path: ./release/*
|
path: ./release/*
|
||||||
@@ -128,21 +138,21 @@ jobs:
|
|||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
uses: actions/download-artifact@v7
|
uses: actions/download-artifact@v8
|
||||||
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@v3
|
uses: docker/login-action@v4
|
||||||
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@v3
|
uses: docker/setup-buildx-action@v4
|
||||||
|
|
||||||
- name: Build and push images
|
- name: Build and push images
|
||||||
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
if: ${{ env.HAS_DOCKER_CREDS == 'true' }}
|
||||||
@@ -163,7 +173,7 @@ jobs:
|
|||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v7
|
uses: actions/download-artifact@v8
|
||||||
with:
|
with:
|
||||||
path: artifacts
|
path: artifacts
|
||||||
|
|
||||||
|
|||||||
@@ -14,10 +14,18 @@ on:
|
|||||||
- 'go.sum'
|
- 'go.sum'
|
||||||
jobs:
|
jobs:
|
||||||
|
|
||||||
smoke-extra:
|
smoke-extra-libvirt:
|
||||||
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: Run extra smoke tests
|
name: ${{ matrix.target }}
|
||||||
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:
|
||||||
@@ -40,28 +48,85 @@ 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: freebsd-amd64
|
- name: ${{ matrix.target }}
|
||||||
run: make smoke-vagrant/freebsd-amd64
|
run: make smoke-vagrant/${{ matrix.target }}
|
||||||
|
|
||||||
- name: openbsd-amd64
|
timeout-minutes: 30
|
||||||
run: make smoke-vagrant/openbsd-amd64
|
|
||||||
|
|
||||||
- name: netbsd-amd64
|
# linux-386 needs VirtualBox, which conflicts with KVM/libvirt -- isolated job.
|
||||||
run: make smoke-vagrant/netbsd-amd64
|
smoke-extra-virtualbox:
|
||||||
|
if: github.ref == 'refs/heads/master' || contains(github.event.pull_request.labels.*.name, 'smoke-test-extra')
|
||||||
|
name: linux-386
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
env:
|
||||||
|
VAGRANT_DEFAULT_PROVIDER: virtualbox
|
||||||
|
steps:
|
||||||
|
|
||||||
- name: linux-amd64-ipv6disable
|
- uses: actions/checkout@v6
|
||||||
run: make smoke-vagrant/linux-amd64-ipv6disable
|
|
||||||
|
|
||||||
# linux-386 runs last because it requires disabling KVM to use VirtualBox,
|
- uses: actions/setup-go@v6
|
||||||
# which prevents libvirt (used by the other tests) from working after this point.
|
with:
|
||||||
- name: install virtualbox for i386 test
|
go-version: '1.25'
|
||||||
|
check-latest: true
|
||||||
|
|
||||||
|
- name: add hashicorp source
|
||||||
|
run: wget -O- https://apt.releases.hashicorp.com/gpg | gpg --dearmor | sudo tee /usr/share/keyrings/hashicorp-archive-keyring.gpg && echo "deb [signed-by=/usr/share/keyrings/hashicorp-archive-keyring.gpg] https://apt.releases.hashicorp.com $(lsb_release -cs) main" | sudo tee /etc/apt/sources.list.d/hashicorp.list
|
||||||
|
|
||||||
|
- name: install vagrant and virtualbox
|
||||||
run: |
|
run: |
|
||||||
sudo apt-get install -y virtualbox
|
sudo apt-get update && sudo apt-get install -y vagrant 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
|
||||||
|
|||||||
@@ -0,0 +1,272 @@
|
|||||||
|
#!/usr/bin/env pwsh
|
||||||
|
# Windows smoke test for the nebula tun + UDP + NLM code paths.
|
||||||
|
#
|
||||||
|
# Topology:
|
||||||
|
# - lighthouse runs natively on the Windows host (wintun + windows UDP)
|
||||||
|
# - peer runs inside WSL2 (Linux build of nebula, /dev/net/tun)
|
||||||
|
#
|
||||||
|
# WSL2 gives us a real netns boundary so the loopback fast-path on Windows
|
||||||
|
# does not short-circuit the overlay -- when WSL pings the lighthouse VPN IP,
|
||||||
|
# Linux has no idea that IP is local to the Windows host, so the packet is
|
||||||
|
# forced through nebula. Same in reverse.
|
||||||
|
|
||||||
|
$ErrorActionPreference = 'Stop'
|
||||||
|
|
||||||
|
# wsl.exe emits UTF-16 LE by default which PowerShell reads as bytes, mangling
|
||||||
|
# every captured string. WSL_UTF8 makes wsl.exe emit UTF-8 instead.
|
||||||
|
$env:WSL_UTF8 = '1'
|
||||||
|
|
||||||
|
$RepoRoot = Resolve-Path "$PSScriptRoot\..\..\.."
|
||||||
|
$Nebula = Join-Path $RepoRoot 'nebula.exe'
|
||||||
|
$NebulaCert = Join-Path $RepoRoot 'nebula-cert.exe'
|
||||||
|
$NebulaLinux = Join-Path $RepoRoot 'build\linux-amd64\nebula'
|
||||||
|
|
||||||
|
if (-not (Test-Path $Nebula)) { throw "missing $Nebula; run 'make bin-windows' first" }
|
||||||
|
if (-not (Test-Path $NebulaCert)) { throw "missing $NebulaCert; run 'make bin-windows' first" }
|
||||||
|
if (-not (Test-Path $NebulaLinux)) { throw "missing $NebulaLinux; build the linux nebula first" }
|
||||||
|
|
||||||
|
# Matches the distro installed by Vampire/setup-wsl in smoke-extra.yml.
|
||||||
|
$Distro = 'Ubuntu-24.04'
|
||||||
|
$listed = (wsl --list --quiet 2>$null) -join "`n"
|
||||||
|
if ($listed -notmatch [regex]::Escape($Distro)) {
|
||||||
|
throw "WSL distro $Distro not registered. Got: $listed"
|
||||||
|
}
|
||||||
|
Write-Host "Using WSL distro: $Distro"
|
||||||
|
|
||||||
|
# Windows host as seen from inside WSL: WSL's default-route gateway. We extract
|
||||||
|
# it with a regex rather than awk fields so PowerShell does not eat any '$N'
|
||||||
|
# tokens, and tabs/double-spaces in `ip route` output do not confuse a cut.
|
||||||
|
$ipCmd = 'ip route show default | grep -oE "([0-9]+\.){3}[0-9]+" | head -1'
|
||||||
|
$WindowsIp = (wsl -d $Distro -- bash -c $ipCmd).Trim()
|
||||||
|
if (-not $WindowsIp) { throw "could not determine Windows host IP from WSL" }
|
||||||
|
Write-Host "Windows host IP from WSL: $WindowsIp"
|
||||||
|
|
||||||
|
$WorkDir = Join-Path $env:TEMP 'nebula-smoke-windows'
|
||||||
|
if (Test-Path $WorkDir) { Remove-Item -Recurse -Force $WorkDir }
|
||||||
|
New-Item -ItemType Directory -Path $WorkDir | Out-Null
|
||||||
|
|
||||||
|
$WslDir = '/tmp/nebula-smoke'
|
||||||
|
wsl -d $Distro -- bash -c "rm -rf $WslDir && mkdir -p $WslDir" | Out-Null
|
||||||
|
|
||||||
|
$DevName = 'nebula-smoke'
|
||||||
|
$Ip1 = '192.168.241.1'
|
||||||
|
$Ip2 = '192.168.241.2'
|
||||||
|
$Port = 4242
|
||||||
|
|
||||||
|
& $NebulaCert ca -name 'smoke-ca' -out-crt "$WorkDir\ca.crt" -out-key "$WorkDir\ca.key"
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "nebula-cert ca failed (exit $LASTEXITCODE)" }
|
||||||
|
|
||||||
|
& $NebulaCert sign -name 'lighthouse' -networks "$Ip1/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\lighthouse.crt" -out-key "$WorkDir\lighthouse.key"
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign lighthouse failed (exit $LASTEXITCODE)" }
|
||||||
|
|
||||||
|
& $NebulaCert sign -name 'peer' -networks "$Ip2/24" -ca-crt "$WorkDir\ca.crt" -ca-key "$WorkDir\ca.key" -out-crt "$WorkDir\peer.crt" -out-key "$WorkDir\peer.key"
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "nebula-cert sign peer failed (exit $LASTEXITCODE)" }
|
||||||
|
|
||||||
|
# Windows lighthouse config.
|
||||||
|
@"
|
||||||
|
pki:
|
||||||
|
ca: $WorkDir\ca.crt
|
||||||
|
cert: $WorkDir\lighthouse.crt
|
||||||
|
key: $WorkDir\lighthouse.key
|
||||||
|
static_host_map: {}
|
||||||
|
lighthouse:
|
||||||
|
am_lighthouse: true
|
||||||
|
interval: 60
|
||||||
|
hosts: []
|
||||||
|
listen:
|
||||||
|
host: 0.0.0.0
|
||||||
|
port: $Port
|
||||||
|
tun:
|
||||||
|
disabled: false
|
||||||
|
dev: $DevName
|
||||||
|
drop_local_broadcast: false
|
||||||
|
drop_multicast: false
|
||||||
|
tx_queue: 500
|
||||||
|
mtu: 1300
|
||||||
|
network_category: private
|
||||||
|
logging:
|
||||||
|
level: info
|
||||||
|
format: text
|
||||||
|
firewall:
|
||||||
|
outbound_action: drop
|
||||||
|
inbound_action: drop
|
||||||
|
conntrack:
|
||||||
|
tcp_timeout: 12m
|
||||||
|
udp_timeout: 3m
|
||||||
|
default_timeout: 10m
|
||||||
|
outbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
inbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
"@ | Out-File -FilePath "$WorkDir\lighthouse.yml" -Encoding utf8
|
||||||
|
|
||||||
|
# WSL peer config (paths are POSIX, deliberately).
|
||||||
|
@"
|
||||||
|
pki:
|
||||||
|
ca: $WslDir/ca.crt
|
||||||
|
cert: $WslDir/peer.crt
|
||||||
|
key: $WslDir/peer.key
|
||||||
|
static_host_map:
|
||||||
|
"${Ip1}": ["${WindowsIp}:$Port"]
|
||||||
|
lighthouse:
|
||||||
|
am_lighthouse: false
|
||||||
|
interval: 60
|
||||||
|
hosts:
|
||||||
|
- "${Ip1}"
|
||||||
|
listen:
|
||||||
|
host: 0.0.0.0
|
||||||
|
port: 0
|
||||||
|
tun:
|
||||||
|
disabled: false
|
||||||
|
dev: nebula1
|
||||||
|
drop_local_broadcast: false
|
||||||
|
drop_multicast: false
|
||||||
|
tx_queue: 500
|
||||||
|
mtu: 1300
|
||||||
|
logging:
|
||||||
|
level: info
|
||||||
|
format: text
|
||||||
|
firewall:
|
||||||
|
outbound_action: drop
|
||||||
|
inbound_action: drop
|
||||||
|
conntrack:
|
||||||
|
tcp_timeout: 12m
|
||||||
|
udp_timeout: 3m
|
||||||
|
default_timeout: 10m
|
||||||
|
outbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
inbound:
|
||||||
|
- port: any
|
||||||
|
proto: any
|
||||||
|
host: any
|
||||||
|
"@ | Out-File -FilePath "$WorkDir\peer.yml" -Encoding utf8
|
||||||
|
|
||||||
|
# Stage WSL artifacts. Convert Windows paths to WSL paths ourselves rather than
|
||||||
|
# calling `wslpath`, because PowerShell's argument-passing to external EXEs
|
||||||
|
# strips backslashes from path arguments in ways that are hard to escape around.
|
||||||
|
function ConvertTo-WslPath {
|
||||||
|
param([string]$WindowsPath)
|
||||||
|
if ($WindowsPath -notmatch '^([A-Za-z]):\\(.*)$') {
|
||||||
|
throw "cannot convert path to WSL: $WindowsPath"
|
||||||
|
}
|
||||||
|
return "/mnt/$($matches[1].ToLower())/$($matches[2].Replace('\','/'))"
|
||||||
|
}
|
||||||
|
|
||||||
|
$WslWorkDir = ConvertTo-WslPath $WorkDir
|
||||||
|
$WslNebulaPath = ConvertTo-WslPath $NebulaLinux
|
||||||
|
wsl -d $Distro -- bash -c "cp '$WslWorkDir/ca.crt' '$WslWorkDir/peer.crt' '$WslWorkDir/peer.key' '$WslWorkDir/peer.yml' $WslDir/ && cp '$WslNebulaPath' $WslDir/nebula && chmod +x $WslDir/nebula"
|
||||||
|
|
||||||
|
# Make sure WSL has tun support and /dev/net/tun is usable before starting
|
||||||
|
# nebula. Diagnostics first so a fail here points at the real problem (e.g.
|
||||||
|
# WSL1 distros do not have a real kernel and will not have tun).
|
||||||
|
Write-Host '=== WSL diagnostic ==='
|
||||||
|
wsl --version 2>&1 | Out-Host
|
||||||
|
wsl --list --verbose 2>&1 | Out-Host
|
||||||
|
wsl -d $Distro -u root -- uname -a | Out-Host
|
||||||
|
wsl -d $Distro -u root -- bash -c "modprobe tun 2>&1 || true; mkdir -p /dev/net; [ -c /dev/net/tun ] || mknod /dev/net/tun c 10 200; chmod 600 /dev/net/tun; ls -l /dev/net/tun"
|
||||||
|
if ($LASTEXITCODE -ne 0) { throw "failed to prepare /dev/net/tun in WSL (TUN support missing?)" }
|
||||||
|
|
||||||
|
# Deliberately no New-NetFirewallRule calls here -- nebula's windows_bypass_wdf
|
||||||
|
# feature is supposed to install WFP permit filters that let inbound traffic
|
||||||
|
# through Windows Defender Firewall on its own. If this smoke regresses, that
|
||||||
|
# feature regressed.
|
||||||
|
|
||||||
|
$lhOut = Join-Path $WorkDir 'lighthouse.out.log'
|
||||||
|
$lhErr = Join-Path $WorkDir 'lighthouse.err.log'
|
||||||
|
$lhProc = Start-Process -FilePath $Nebula -ArgumentList @('-config', "$WorkDir\lighthouse.yml") `
|
||||||
|
-PassThru -NoNewWindow `
|
||||||
|
-RedirectStandardOutput $lhOut `
|
||||||
|
-RedirectStandardError $lhErr
|
||||||
|
|
||||||
|
# Run nebula in WSL as root with no sudo + no shell wrapper. PowerShell's
|
||||||
|
# Start-Process arg quoting mangles `bash -c "..."` strings that contain
|
||||||
|
# spaces/redirections, so we skip bash entirely and let Start-Process do the
|
||||||
|
# stdout/stderr capture itself.
|
||||||
|
$peerOut = Join-Path $WorkDir 'peer.out.log'
|
||||||
|
$peerErr = Join-Path $WorkDir 'peer.err.log'
|
||||||
|
$peerProc = Start-Process -FilePath 'wsl' `
|
||||||
|
-ArgumentList @('-d', $Distro, '-u', 'root', '--', "$WslDir/nebula", '-config', "$WslDir/peer.yml") `
|
||||||
|
-PassThru -NoNewWindow `
|
||||||
|
-RedirectStandardOutput $peerOut `
|
||||||
|
-RedirectStandardError $peerErr
|
||||||
|
|
||||||
|
function Wait-Until {
|
||||||
|
param([scriptblock]$Predicate, [int]$TimeoutSec, [string]$What)
|
||||||
|
$deadline = (Get-Date).AddSeconds($TimeoutSec)
|
||||||
|
while ((Get-Date) -lt $deadline) {
|
||||||
|
if (& $Predicate) { return }
|
||||||
|
Start-Sleep -Milliseconds 500
|
||||||
|
}
|
||||||
|
throw "timed out waiting for: $What"
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
Wait-Until -TimeoutSec 30 -What "windows wintun adapter $DevName with NetworkCategory=Private" -Predicate {
|
||||||
|
if ($lhProc.HasExited) { throw "lighthouse exited (code $($lhProc.ExitCode)) before tun was ready" }
|
||||||
|
$p = Get-NetConnectionProfile -InterfaceAlias $DevName -ErrorAction SilentlyContinue
|
||||||
|
$p -and ("$($p.NetworkCategory)" -ieq 'Private')
|
||||||
|
}
|
||||||
|
Write-Host "OK: $DevName NetworkCategory=Private"
|
||||||
|
|
||||||
|
Wait-Until -TimeoutSec 30 -What "WSL nebula1 with $Ip2" -Predicate {
|
||||||
|
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before tun was ready" }
|
||||||
|
$r = wsl -d $Distro -u root -- bash -c "ip -o addr show nebula1 2>/dev/null | grep -q 'inet $Ip2' && echo yes"
|
||||||
|
("$r").Trim() -eq 'yes'
|
||||||
|
}
|
||||||
|
Write-Host "OK: WSL nebula1 has $Ip2"
|
||||||
|
|
||||||
|
Wait-Until -TimeoutSec 30 -What "ping from WSL peer to windows lighthouse ($Ip1)" -Predicate {
|
||||||
|
if ($peerProc.HasExited) { throw "peer exited (code $($peerProc.ExitCode)) before ping succeeded" }
|
||||||
|
$r = wsl -d $Distro -u root -- bash -c "ping -c1 -W1 $Ip1 >/dev/null 2>&1 && echo OK"
|
||||||
|
("$r").Trim() -eq 'OK'
|
||||||
|
}
|
||||||
|
Write-Host "OK: WSL peer -> windows lighthouse"
|
||||||
|
|
||||||
|
Wait-Until -TimeoutSec 30 -What "ping from windows lighthouse to WSL peer ($Ip2)" -Predicate {
|
||||||
|
$null = & ping.exe -n 1 -w 1000 $Ip2
|
||||||
|
$LASTEXITCODE -eq 0
|
||||||
|
}
|
||||||
|
Write-Host "OK: windows lighthouse -> WSL peer"
|
||||||
|
|
||||||
|
Write-Host ''
|
||||||
|
Write-Host 'All smoke checks passed.'
|
||||||
|
}
|
||||||
|
catch {
|
||||||
|
Write-Host ''
|
||||||
|
Write-Host '=== lighthouse stdout ==='
|
||||||
|
Get-Content $lhOut -ErrorAction SilentlyContinue | Out-Host
|
||||||
|
Write-Host '=== lighthouse stderr ==='
|
||||||
|
Get-Content $lhErr -ErrorAction SilentlyContinue | Out-Host
|
||||||
|
Write-Host '=== peer stdout ==='
|
||||||
|
Get-Content $peerOut -ErrorAction SilentlyContinue | Out-Host
|
||||||
|
Write-Host '=== peer stderr ==='
|
||||||
|
Get-Content $peerErr -ErrorAction SilentlyContinue | Out-Host
|
||||||
|
Write-Host '=== nebula WFP filters ==='
|
||||||
|
# Dump nebula-installed filters so we can verify they got registered with
|
||||||
|
# the conditions we expect.
|
||||||
|
$wfpDump = Join-Path $WorkDir 'wfp.xml'
|
||||||
|
netsh wfp show filters file=$wfpDump 2>&1 | Out-Null
|
||||||
|
if (Test-Path $wfpDump) {
|
||||||
|
Select-String -Path $wfpDump -Pattern 'Nebula' -Context 0,80 -ErrorAction SilentlyContinue | Out-Host
|
||||||
|
}
|
||||||
|
throw
|
||||||
|
}
|
||||||
|
finally {
|
||||||
|
if (-not $lhProc.HasExited) {
|
||||||
|
Stop-Process -Id $lhProc.Id -Force -ErrorAction SilentlyContinue
|
||||||
|
$lhProc.WaitForExit(5000) | Out-Null
|
||||||
|
}
|
||||||
|
wsl -d $Distro -u root -- bash -c "pkill -f $WslDir/nebula 2>/dev/null; true" | Out-Null
|
||||||
|
# pkill returns 1 when no match and wsl propagates that; the smoke is done
|
||||||
|
# so we don't want it to leak into the script's exit code.
|
||||||
|
$global:LASTEXITCODE = 0
|
||||||
|
if ($peerProc -and -not $peerProc.HasExited) {
|
||||||
|
Stop-Process -Id $peerProc.Id -Force -ErrorAction SilentlyContinue
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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 = "generic/netbsd9"
|
config.vm.box = "DefinedNet/netbsd10"
|
||||||
|
|
||||||
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@v6
|
- uses: actions/upload-artifact@v7
|
||||||
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@v6
|
- uses: actions/upload-artifact@v7
|
||||||
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,24 +2,42 @@ 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 []bool
|
bits []uint64
|
||||||
lostCounter metrics.Counter
|
lostCounter metrics.Counter
|
||||||
dupeCounter metrics.Counter
|
dupeCounter metrics.Counter
|
||||||
outOfWindowCounter metrics.Counter
|
outOfWindowCounter metrics.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewBits(bits uint64) *Bits {
|
func NewBits(length uint64) *Bits {
|
||||||
|
if length == 0 || length&(length-1) != 0 {
|
||||||
|
panic(fmt.Sprintf("Bits length must be a power of two, got %d", length))
|
||||||
|
}
|
||||||
|
|
||||||
|
nWords := length / bitsPerWord
|
||||||
|
if nWords == 0 {
|
||||||
|
nWords = 1
|
||||||
|
}
|
||||||
b := &Bits{
|
b := &Bits{
|
||||||
length: bits,
|
length: length,
|
||||||
bits: make([]bool, bits, bits),
|
lengthMask: length - 1,
|
||||||
|
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),
|
||||||
@@ -27,71 +45,194 @@ func NewBits(bits 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] = true
|
b.bits[0] = 1
|
||||||
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 i is within the window, check if it's been set already.
|
if b.strictlyWithinWindow(i) {
|
||||||
if i > b.current-b.length || i < b.length && b.current < b.length {
|
return !b.get(i)
|
||||||
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)",
|
l.Debug("rejected a packet (top)", "current", b.current, "incoming", i)
|
||||||
"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 {
|
||||||
// If i is the next number, return true and update current.
|
// Fast path: i is the next expected counter. Split out so the function
|
||||||
|
// stays small and avoids paying for the slow paths' slog argument-build
|
||||||
|
// stack frame on every call. The bit read/test/write is inlined to
|
||||||
|
// touch the backing word once.
|
||||||
if i == b.current+1 {
|
if i == b.current+1 {
|
||||||
// Check if the oldest bit was lost since we are shifting the window by 1 and occupying it with this counter
|
pos := i & b.lengthMask
|
||||||
// The very first window can only be tracked as lost once we are on the 2nd window or greater
|
word := pos >> 6
|
||||||
if b.bits[i%b.length] == false && i > b.length {
|
mask := uint64(1) << (pos & 63)
|
||||||
|
w := b.bits[word]
|
||||||
|
if i > b.length && w&mask == 0 {
|
||||||
b.lostCounter.Inc(1)
|
b.lostCounter.Inc(1)
|
||||||
}
|
}
|
||||||
b.bits[i%b.length] = true
|
b.bits[word] = w | mask
|
||||||
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 {
|
||||||
lost := int64(0)
|
end := i
|
||||||
// Zero out the bits between the current and the new counter value, limited by the window size,
|
if end > b.current+b.length {
|
||||||
// since the window is shifting
|
end = b.current + b.length
|
||||||
for n := b.current + 1; n <= min(i, b.current+b.length); n++ {
|
}
|
||||||
if b.bits[n%b.length] == false && n > b.length {
|
count := end - b.current
|
||||||
lost++
|
startPos := (b.current + 1) & b.lengthMask
|
||||||
|
|
||||||
|
var lost int64
|
||||||
|
if b.current >= b.length {
|
||||||
|
// Steady state: every cleared slot is past warmup, so any unset
|
||||||
|
// bit we evict is a lost packet from the previous cycle.
|
||||||
|
wasSet := b.clearRange(startPos, count)
|
||||||
|
lost = int64(count) - int64(wasSet)
|
||||||
|
} else {
|
||||||
|
// Warmup (the very first window). Some cleared slots represent
|
||||||
|
// packets <= length where eviction is not "lost" in the usual
|
||||||
|
// sense. This branch is taken at most once per connection so we
|
||||||
|
// don't bother optimizing it.
|
||||||
|
for n := b.current + 1; n <= end; n++ {
|
||||||
|
if !b.get(n) && n > b.length {
|
||||||
|
lost++
|
||||||
|
}
|
||||||
}
|
}
|
||||||
b.bits[n%b.length] = false
|
b.clearRange(startPos, count)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Only record any skipped packets as a result of the window moving further than the window length
|
// Anything past the new window can never be backfilled, so it's lost.
|
||||||
// Any loss within the new window will be accounted for in future calls
|
if i > b.current+b.length {
|
||||||
lost += max(0, int64(i-b.current-b.length))
|
lost += int64(i - b.current - b.length)
|
||||||
|
}
|
||||||
b.lostCounter.Inc(lost)
|
b.lostCounter.Inc(lost)
|
||||||
|
|
||||||
b.bits[i%b.length] = true
|
b.set(i)
|
||||||
b.current = i
|
b.current = i
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
// If i is within the current window but below the current counter,
|
// If i is within the current window but below the current counter, check to see if it's a duplicate
|
||||||
// Check to see if it's a duplicate
|
if b.strictlyWithinWindow(i) {
|
||||||
if i > b.current-b.length || i < b.length && b.current < b.length {
|
pos := i & b.lengthMask
|
||||||
if b.current == i || b.bits[i%b.length] == true {
|
word := pos >> 6
|
||||||
|
mask := uint64(1) << (pos & 63)
|
||||||
|
w := b.bits[word]
|
||||||
|
if b.current == i || w&mask != 0 {
|
||||||
if l.Enabled(context.Background(), slog.LevelDebug) {
|
if l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
l.Debug("Receive window",
|
l.Debug("Receive window",
|
||||||
"accepted", false,
|
"accepted", false,
|
||||||
@@ -104,7 +245,7 @@ func (b *Bits) Update(l *slog.Logger, i uint64) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
b.bits[i%b.length] = true
|
b.bits[word] = w | mask
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+277
-130
@@ -7,61 +7,79 @@ 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(10)
|
b := NewBits(16)
|
||||||
|
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}
|
g := []bool{true, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.snapshot())
|
||||||
|
|
||||||
// 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}
|
g = []bool{true, true, true, false, false, false, false, false, false, false, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.snapshot())
|
||||||
|
|
||||||
// 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 15, which should clear everything and set the 6th element
|
// Jump ahead to 25, which clears the window and sets slot 25%16 = 9.
|
||||||
assert.True(t, b.Check(l, 15))
|
assert.True(t, b.Check(l, 25))
|
||||||
assert.True(t, b.Update(l, 15))
|
assert.True(t, b.Update(l, 25))
|
||||||
assert.EqualValues(t, 15, b.current)
|
assert.EqualValues(t, 25, b.current)
|
||||||
g = []bool{false, false, false, false, false, true, false, false, false, false}
|
g = []bool{false, false, false, false, false, false, false, false, false, true, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.snapshot())
|
||||||
|
|
||||||
// Mark 14, which is allowed because it is in the window
|
// Mark 24, which is in window (current 25, length 16, window covers [10,25]).
|
||||||
assert.True(t, b.Check(l, 14))
|
assert.True(t, b.Check(l, 24))
|
||||||
assert.True(t, b.Update(l, 14))
|
assert.True(t, b.Update(l, 24))
|
||||||
assert.EqualValues(t, 15, b.current)
|
assert.EqualValues(t, 25, b.current)
|
||||||
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.snapshot())
|
||||||
|
|
||||||
// Mark 5, which is not allowed because it is not in the window
|
// Mark 5, not allowed because 5 <= current-length (25-16=9).
|
||||||
assert.False(t, b.Check(l, 5))
|
assert.False(t, b.Check(l, 5))
|
||||||
assert.False(t, b.Update(l, 5))
|
assert.False(t, b.Update(l, 5))
|
||||||
assert.EqualValues(t, 15, b.current)
|
assert.EqualValues(t, 25, b.current)
|
||||||
g = []bool{false, false, false, false, true, true, false, false, false, false}
|
g = []bool{false, false, false, false, false, false, false, false, true, true, false, false, false, false, false, false}
|
||||||
assert.Equal(t, g, b.bits)
|
assert.Equal(t, g, b.snapshot())
|
||||||
|
|
||||||
// make sure we handle wrapping around once to the current position
|
// Make sure we handle wrapping around once to the same slot. With
|
||||||
b = NewBits(10)
|
// length=16, packets 1 and 17 share slot 1.
|
||||||
|
b = NewBits(16)
|
||||||
assert.True(t, b.Update(l, 1))
|
assert.True(t, b.Update(l, 1))
|
||||||
assert.True(t, b.Update(l, 11))
|
assert.True(t, b.Update(l, 17))
|
||||||
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false}, b.bits)
|
assert.Equal(t, []bool{false, true, false, false, false, false, false, false, false, false, false, false, false, false, false, false}, b.snapshot())
|
||||||
|
|
||||||
// Walk through a few windows in order
|
// Walk through a few windows in order
|
||||||
b = NewBits(10)
|
b = NewBits(16)
|
||||||
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)
|
||||||
@@ -72,24 +90,31 @@ 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())
|
||||||
|
|
||||||
b = NewBits(10)
|
// Update(100): clears 16 slots starting at slot 56%16=8. Only slot 7 (for
|
||||||
b.lostCounter.Clear()
|
// packet 55) was set, so 16 - 1 = 15 evicted slots had unset bits.
|
||||||
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
|
// Plus 100 - 55 - 16 = 29 packets fell past the window. Total 44.
|
||||||
assert.Equal(t, int64(45), b.lostCounter.Count())
|
assert.True(t, b.Update(l, 100))
|
||||||
|
assert.Equal(t, int64(39+44), b.lostCounter.Count())
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 100)) // We saw packet 55 and 100 and can still track 90,91,92,93,94,95,96,97,98,99
|
// Update(200): same shape: 16 - 1 = 15 evicted unset, plus 200 - 100 - 16 = 84 past window. Total 99.
|
||||||
assert.Equal(t, int64(89), b.lostCounter.Count())
|
assert.True(t, b.Update(l, 200))
|
||||||
|
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(10)
|
b := NewBits(16)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
@@ -114,120 +139,117 @@ func TestBitsDupeCounter(t *testing.T) {
|
|||||||
|
|
||||||
func TestBitsOutOfWindowCounter(t *testing.T) {
|
func TestBitsOutOfWindowCounter(t *testing.T) {
|
||||||
l := test.NewLogger()
|
l := test.NewLogger()
|
||||||
b := NewBits(10)
|
b := NewBits(16)
|
||||||
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())
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 21))
|
// 9 single-step advances, each evicts a slot whose bit was cleared during
|
||||||
assert.True(t, b.Update(l, 22))
|
// the jump above and whose value was never seen, so each contributes 1
|
||||||
assert.True(t, b.Update(l, 23))
|
// to lostCounter.
|
||||||
assert.True(t, b.Update(l, 24))
|
for n := uint64(21); n <= 29; n++ {
|
||||||
assert.True(t, b.Update(l, 25))
|
assert.True(t, b.Update(l, n))
|
||||||
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())
|
||||||
|
|
||||||
assert.Equal(t, int64(19), b.lostCounter.Count()) // packet 0 wasn't lost
|
// 4 from the Update(20) jump + 9 from 21..29.
|
||||||
|
assert.Equal(t, int64(13), b.lostCounter.Count())
|
||||||
assert.Equal(t, int64(0), b.dupeCounter.Count())
|
assert.Equal(t, int64(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(10)
|
b := NewBits(16)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 20))
|
// Walk 20..29 like the original, just with a bigger window. Same
|
||||||
assert.True(t, b.Update(l, 21))
|
// reasoning as TestBitsOutOfWindowCounter: 4 past-window from Update(20),
|
||||||
assert.True(t, b.Update(l, 22))
|
// then 9 more from the unit advances.
|
||||||
assert.True(t, b.Update(l, 23))
|
for n := uint64(20); 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.Equal(t, int64(13), b.lostCounter.Count())
|
||||||
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(10)
|
b = NewBits(16)
|
||||||
b.lostCounter.Clear()
|
b.lostCounter.Clear()
|
||||||
b.dupeCounter.Clear()
|
b.dupeCounter.Clear()
|
||||||
b.outOfWindowCounter.Clear()
|
b.outOfWindowCounter.Clear()
|
||||||
|
|
||||||
assert.True(t, b.Update(l, 9))
|
// Update(15) clears the warmup window (no lost), sets slot 15.
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
// 10 will set 0 index, 0 was already set, no lost packets
|
|
||||||
assert.True(t, b.Update(l, 10))
|
|
||||||
assert.Equal(t, int64(0), b.lostCounter.Count())
|
|
||||||
// 11 will set 1 index, 1 was missed, we should see 1 packet lost
|
|
||||||
assert.True(t, b.Update(l, 11))
|
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
|
||||||
// Now let's fill in the window, should end up with 8 lost packets
|
|
||||||
assert.True(t, b.Update(l, 12))
|
|
||||||
assert.True(t, b.Update(l, 13))
|
|
||||||
assert.True(t, b.Update(l, 14))
|
|
||||||
assert.True(t, b.Update(l, 15))
|
assert.True(t, b.Update(l, 15))
|
||||||
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
|
||||||
|
// Update(16): slot 0 was already set (NewBits seeded it), and 16 is not
|
||||||
|
// strictly > length, so nothing is recorded as lost.
|
||||||
assert.True(t, b.Update(l, 16))
|
assert.True(t, b.Update(l, 16))
|
||||||
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
|
||||||
|
// Update(17): we jumped straight from 0 to 15, so slot 1 was cleared
|
||||||
|
// (and never re-set). 17 > 16 is past warmup, so packet 1 is recorded lost.
|
||||||
assert.True(t, b.Update(l, 17))
|
assert.True(t, b.Update(l, 17))
|
||||||
assert.True(t, b.Update(l, 18))
|
assert.Equal(t, int64(1), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 19))
|
|
||||||
assert.Equal(t, int64(8), b.lostCounter.Count())
|
|
||||||
|
|
||||||
// Jump ahead by a window size
|
// Fill in 18..30 in single steps. Each i evicts slot i%16. Slots 2..14
|
||||||
assert.True(t, b.Update(l, 29))
|
// were all cleared during Update(15), and we never re-set any of them,
|
||||||
assert.Equal(t, int64(8), b.lostCounter.Count())
|
// so each i in 18..30 is a fresh lost packet — 13 more.
|
||||||
// Now lets walk ahead normally through the window, the missed packets should fill in
|
for n := uint64(18); n <= 30; n++ {
|
||||||
assert.True(t, b.Update(l, 30))
|
assert.True(t, b.Update(l, n))
|
||||||
assert.True(t, b.Update(l, 31))
|
}
|
||||||
assert.True(t, b.Update(l, 32))
|
assert.Equal(t, int64(14), b.lostCounter.Count())
|
||||||
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 2 windows, should have recording 1 full window missing
|
// Jump ahead by exactly one window size.
|
||||||
assert.True(t, b.Update(l, 58))
|
assert.True(t, b.Update(l, 46))
|
||||||
assert.Equal(t, int64(27), b.lostCounter.Count())
|
// end = min(46, 30+16) = 46, count = 16, all slots cleared. Before the
|
||||||
// Now lets walk ahead normally through the window, the missed packets should fill in from this window
|
// jump every slot 0..15 had been set (Update(15), (16), (17), 18..30),
|
||||||
assert.True(t, b.Update(l, 59))
|
// so wasSet=16 and 46 == current+length means no past-window slack:
|
||||||
assert.True(t, b.Update(l, 60))
|
// lost contribution = 0.
|
||||||
assert.True(t, b.Update(l, 61))
|
assert.Equal(t, int64(14), b.lostCounter.Count())
|
||||||
assert.True(t, b.Update(l, 62))
|
|
||||||
assert.True(t, b.Update(l, 63))
|
// Walk 47..55. The Update(46) jump cleared every slot, so only slot 14
|
||||||
assert.True(t, b.Update(l, 64))
|
// (for packet 46) is set when we start. Each subsequent unit step lands
|
||||||
assert.True(t, b.Update(l, 65))
|
// on a slot that was cleared and is past warmup, so it counts as lost.
|
||||||
assert.True(t, b.Update(l, 66))
|
// 9 more = 23.
|
||||||
assert.True(t, b.Update(l, 67))
|
for n := uint64(47); n <= 55; n++ {
|
||||||
// 68 packets tracked, 32 seen, 36 missed
|
assert.True(t, b.Update(l, n))
|
||||||
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(10)
|
b := NewBits(16)
|
||||||
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))
|
||||||
@@ -244,7 +266,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())
|
||||||
// assert.True(t, b.Update(l, 8))
|
// Skip packet 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))
|
||||||
@@ -252,9 +274,23 @@ 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
|
|
||||||
assert.True(t, b.Update(l, 19))
|
// Jump to 25. With length=16, slot 25%16=9 corresponds to packet 9
|
||||||
|
// (which we DID receive), so its bit is set and no lost++ from that
|
||||||
|
// eviction. The trace below shows the only loss is packet 8.
|
||||||
|
assert.True(t, b.Update(l, 25))
|
||||||
|
// current was 14, i=25. end=min(25,30)=25. count=11. startPos=15.
|
||||||
|
// steady? current=14<16, so warmup branch: per-bit n=15..25, count those
|
||||||
|
// with !get(n) AND n>16. n=17..25 are >16. Among slots 17%16=1..25%16=9
|
||||||
|
// did we set slots 1..9 (packets 1..9)? Yes for all but slot 8 (packet 8
|
||||||
|
// was skipped). n=24 maps to slot 8 which is FALSE → lost++. All other
|
||||||
|
// n in 17..25 map to slots that are set. n=16 is not strictly > 16. So
|
||||||
|
// lost = 1.
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
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))
|
||||||
@@ -263,29 +299,140 @@ 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
|
// We missed packet 8 above and that loss is still recorded once, never
|
||||||
|
// double-counted, never zeroed.
|
||||||
assert.Equal(t, int64(1), b.lostCounter.Count())
|
assert.Equal(t, int64(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())
|
||||||
}
|
}
|
||||||
|
|
||||||
func BenchmarkBits(b *testing.B) {
|
// TestBitsWarmupOvershoot exercises the jump path's warmup arm with an
|
||||||
z := NewBits(10)
|
// overshoot past one full window. NewBits leaves current=0 with only slot 0
|
||||||
for n := 0; n < b.N; n++ {
|
// "set" by the marker. Jumping straight to length+k must (a) clear every
|
||||||
for i := range z.bits {
|
// slot the jump straddles, (b) count only past-window slack (not the
|
||||||
z.bits[i] = true
|
// in-window slots, which never had a "lost" tenant during warmup), and
|
||||||
}
|
// (c) leave the cursor at the new counter so subsequent unit advances
|
||||||
for i := range z.bits {
|
// count from steady state. The marker bit at slot 0 is irrelevant once
|
||||||
z.bits[i] = false
|
// current >= length.
|
||||||
}
|
func TestBitsWarmupOvershoot(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
b := NewBits(16)
|
||||||
|
b.lostCounter.Clear()
|
||||||
|
|
||||||
|
// Jump from current=0 to i=20 (length=16, overshoot=4).
|
||||||
|
// Warmup arm: counts slots in [1..16] where bit unset and n>length.
|
||||||
|
// Only n=16 was unset and >length: but slot 16%16=0 is the marker,
|
||||||
|
// so b.get(16) reads bits[0]=1 and skips. Result: 0 lost from the loop.
|
||||||
|
// Past-window: i - current - length = 20 - 0 - 16 = 4 lost.
|
||||||
|
assert.True(t, b.Update(l, 20))
|
||||||
|
assert.Equal(t, int64(4), b.lostCounter.Count())
|
||||||
|
assert.Equal(t, uint64(20), b.current)
|
||||||
|
|
||||||
|
// Steady state now (current=20 >= length=16). Unit advance to 21
|
||||||
|
// stomps slot 21%16=5, which was cleared by the jump and not reset,
|
||||||
|
// so this is +1 lost.
|
||||||
|
assert.True(t, b.Update(l, 21))
|
||||||
|
assert.Equal(t, int64(5), b.lostCounter.Count())
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBitsCheckAcrossWarmupBoundary pins the underflow trick in Check's
|
||||||
|
// in-window clause. While in warmup, b.current-b.length underflows uint64
|
||||||
|
// to a huge value so the first OR-clause is always false; the second
|
||||||
|
// clause (i < length && current < length) carries the in-window check.
|
||||||
|
// Once current >= length the regimes flip cleanly.
|
||||||
|
func TestBitsCheckAcrossWarmupBoundary(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
b := NewBits(16)
|
||||||
|
|
||||||
|
// Warmup: current=0. Check(0) must read the marker (set) and return false.
|
||||||
|
assert.False(t, b.Check(l, 0), "marker slot should look already-received")
|
||||||
|
// Warmup: any 0 < i < length is in-window and unset → accepted.
|
||||||
|
for i := uint64(1); i < 16; i++ {
|
||||||
|
assert.True(t, b.Check(l, i), "warmup in-window i=%d should be accepted", i)
|
||||||
|
}
|
||||||
|
// Warmup: i >= length but > current is "next number" so accepted.
|
||||||
|
assert.True(t, b.Check(l, 16))
|
||||||
|
assert.True(t, b.Check(l, 1_000_000))
|
||||||
|
|
||||||
|
// Cross into steady state.
|
||||||
|
assert.True(t, b.Update(l, 100))
|
||||||
|
// Now current=100, length=16. In-window range is [85..100].
|
||||||
|
// 84 is just outside: the underflow clause activates; 84 > 100-16=84 is false.
|
||||||
|
// And the warmup clause is false (current >= length). So out of window.
|
||||||
|
assert.False(t, b.Check(l, 84))
|
||||||
|
// 85 sits at the boundary. 85 > 84 is true → in window, unset → accept.
|
||||||
|
assert.True(t, b.Check(l, 85))
|
||||||
|
// 100 is current itself; not strictly greater, in-window, but already set.
|
||||||
|
assert.False(t, b.Check(l, 100))
|
||||||
|
// Way out: clearly out of window.
|
||||||
|
assert.False(t, b.Check(l, 50))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBitsMarkerInvariant verifies the seeded bits[0]=1 marker behaves
|
||||||
|
// correctly across warmup and beyond. Update should never clear the marker
|
||||||
|
// during warmup (clearRange skips position 0 when startPos=1), and once
|
||||||
|
// current >= length the marker is no longer consulted by Check/Update on
|
||||||
|
// the live path — but it must still report counter 0 as a duplicate while
|
||||||
|
// we are in warmup.
|
||||||
|
func TestBitsMarkerInvariant(t *testing.T) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
b := NewBits(8)
|
||||||
|
|
||||||
|
// Counter 0 is the seeded marker; Check sees it as already received.
|
||||||
|
assert.False(t, b.Check(l, 0))
|
||||||
|
// Update(0) at current=0 hits the duplicate branch.
|
||||||
|
b.dupeCounter.Clear()
|
||||||
|
assert.False(t, b.Update(l, 0))
|
||||||
|
assert.Equal(t, int64(1), b.dupeCounter.Count())
|
||||||
|
|
||||||
|
// Walk forward through warmup; the marker must remain set.
|
||||||
|
for n := uint64(1); n <= 7; n++ {
|
||||||
|
assert.True(t, b.Update(l, n))
|
||||||
|
}
|
||||||
|
// Position 0 (the marker) should still read as set because we never
|
||||||
|
// cleared it; Update(0) still looks like a duplicate.
|
||||||
|
assert.False(t, b.Check(l, 0))
|
||||||
|
|
||||||
|
// Cross into steady state with a unit advance to 8: pos=0, evicts the
|
||||||
|
// marker bit. The lost-counter guard (i > b.length) is false (8 == 8),
|
||||||
|
// so this advance does NOT charge a lost packet — exactly what the
|
||||||
|
// marker is there to prevent.
|
||||||
|
b.lostCounter.Clear()
|
||||||
|
assert.True(t, b.Update(l, 8))
|
||||||
|
assert.Equal(t, int64(0), b.lostCounter.Count())
|
||||||
|
// The slot at pos 0 is now occupied by counter 8.
|
||||||
|
assert.False(t, b.Check(l, 8))
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkBitsUpdateInOrder is the steady-state hot path: each call is
|
||||||
|
// i == current+1.
|
||||||
|
func BenchmarkBitsUpdateInOrder(b *testing.B) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
z := NewBits(16384)
|
||||||
|
for n := 0; n < b.N; n++ {
|
||||||
|
z.Update(l, uint64(n)+1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkBitsUpdateReorder simulates light reorder within the window:
|
||||||
|
// every other packet arrives one slot behind its predecessor (forces the
|
||||||
|
// in-window backfill branch).
|
||||||
|
func BenchmarkBitsUpdateReorder(b *testing.B) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
z := NewBits(16384)
|
||||||
|
for n := 0; n < b.N; n++ {
|
||||||
|
base := uint64(n) * 2
|
||||||
|
z.Update(l, base+2)
|
||||||
|
z.Update(l, base+1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkBitsUpdateLargeJumps stresses the clearRange word-level path.
|
||||||
|
func BenchmarkBitsUpdateLargeJumps(b *testing.B) {
|
||||||
|
l := test.NewLogger()
|
||||||
|
z := NewBits(16384)
|
||||||
|
for n := 0; n < b.N; n++ {
|
||||||
|
z.Update(l, uint64(n+1)*1000)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -217,6 +217,10 @@ func (ncp *CAPool) verify(c Certificate, now time.Time, certFp string, signerFp
|
|||||||
return nil, err
|
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,3 +654,31 @@ 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,6 +112,9 @@ 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,6 +151,9 @@ 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,6 +22,7 @@ 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")
|
||||||
|
|||||||
+12
-42
@@ -11,7 +11,6 @@ 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"
|
||||||
@@ -45,19 +44,16 @@ 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)
|
||||||
@@ -369,7 +365,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.sendPunch(hostinfo)
|
cm.punchy.SendPunch(hostinfo)
|
||||||
}
|
}
|
||||||
|
|
||||||
return decision, hostinfo, primary
|
return decision, hostinfo, primary
|
||||||
@@ -400,17 +396,16 @@ 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.sendPunch(hostinfo)
|
cm.punchy.SendPunch(hostinfo)
|
||||||
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
cm.trafficTimer.Add(hostinfo.localIndexId, cm.checkInterval)
|
||||||
return doNothing, nil, nil
|
return doNothing, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if cm.punchy.GetTargetEverything() {
|
// We aren't receiving traffic but we are sending it. The outbound
|
||||||
// This is similar to the old punchy behavior with a slight optimization.
|
// traffic itself refreshes the primary remote's NAT state; this
|
||||||
// We aren't receiving traffic but we are sending it, punch on all known
|
// fans out to non-primary remotes, but only if target_all_remotes
|
||||||
// ips in case we need to re-prime NAT state
|
// is configured.
|
||||||
cm.sendPunch(hostinfo)
|
cm.punchy.SendPunchToAll(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",
|
||||||
@@ -512,31 +507,6 @@ 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
|
||||||
|
|||||||
@@ -64,7 +64,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)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -146,7 +146,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)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||||
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
nc := newConnectionManagerFromConfig(test.NewLogger(), conf, hostMap, punchy)
|
||||||
nc.intf = ifce
|
nc.intf = ifce
|
||||||
p := []byte("")
|
p := []byte("")
|
||||||
@@ -233,7 +233,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)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||||
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
|
||||||
@@ -358,7 +358,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)
|
punchy := NewPunchyFromConfig(test.NewLogger(), conf, nil)
|
||||||
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
|
||||||
|
|||||||
+5
-4
@@ -7,13 +7,14 @@ import (
|
|||||||
|
|
||||||
"github.com/slackhq/nebula/cert"
|
"github.com/slackhq/nebula/cert"
|
||||||
"github.com/slackhq/nebula/handshake"
|
"github.com/slackhq/nebula/handshake"
|
||||||
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
)
|
)
|
||||||
|
|
||||||
const ReplayWindow = 1024
|
const ReplayWindow = 1024
|
||||||
|
|
||||||
type ConnectionState struct {
|
type ConnectionState struct {
|
||||||
eKey *NebulaCipherState
|
eKey noiseutil.CipherState
|
||||||
dKey *NebulaCipherState
|
dKey noiseutil.CipherState
|
||||||
myCert cert.Certificate
|
myCert cert.Certificate
|
||||||
peerCert *cert.CachedCertificate
|
peerCert *cert.CachedCertificate
|
||||||
initiator bool
|
initiator bool
|
||||||
@@ -31,8 +32,8 @@ func newConnectionStateFromResult(r *handshake.Result) *ConnectionState {
|
|||||||
myCert: r.MyCert,
|
myCert: r.MyCert,
|
||||||
initiator: r.Initiator,
|
initiator: r.Initiator,
|
||||||
peerCert: r.RemoteCert,
|
peerCert: r.RemoteCert,
|
||||||
eKey: NewNebulaCipherState(r.EKey),
|
eKey: noiseutil.NewCipherState(r.EKey, r.Cipher),
|
||||||
dKey: NewNebulaCipherState(r.DKey),
|
dKey: noiseutil.NewCipherState(r.DKey, r.Cipher),
|
||||||
window: NewBits(ReplayWindow),
|
window: NewBits(ReplayWindow),
|
||||||
}
|
}
|
||||||
ci.messageCounter.Add(r.MessageIndex)
|
ci.messageCounter.Add(r.MessageIndex)
|
||||||
|
|||||||
+2
-4
@@ -4,15 +4,13 @@
|
|||||||
package e2e
|
package e2e
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"io"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"dario.cat/mergo"
|
"dario.cat/mergo"
|
||||||
"github.com/google/gopacket"
|
"github.com/google/gopacket"
|
||||||
"github.com/google/gopacket/layers"
|
"github.com/google/gopacket/layers"
|
||||||
@@ -382,7 +380,7 @@ func getAddrs(ns []netip.Prefix) []netip.Addr {
|
|||||||
func NewTestLogger() *slog.Logger {
|
func NewTestLogger() *slog.Logger {
|
||||||
v := os.Getenv("TEST_LOGS")
|
v := os.Getenv("TEST_LOGS")
|
||||||
if v == "" {
|
if v == "" {
|
||||||
return slog.New(slog.NewTextHandler(io.Discard, nil))
|
return slog.New(slog.DiscardHandler)
|
||||||
}
|
}
|
||||||
|
|
||||||
level := slog.LevelInfo
|
level := slog.LevelInfo
|
||||||
|
|||||||
+3
-7
@@ -18,14 +18,10 @@ import (
|
|||||||
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
// retry mechanism gives the wg.Wait()-driven goroutines a moment to drain
|
||||||
// before failing the assertion.
|
// before failing the assertion.
|
||||||
//
|
//
|
||||||
// IgnoreCurrent is necessary in the parallelized suite: other tests can
|
// Intentionally NOT t.Parallel()'d: concurrent tests would have their own
|
||||||
// leave goroutines mid-shutdown when this one runs (Stop is async, the
|
// goroutines running and trip the assertion.
|
||||||
// wg.Wait() drain is not blocking on test return). We're checking that
|
|
||||||
// *this* test's setup tears down cleanly, not that the whole suite is
|
|
||||||
// idle at this moment. Intentionally NOT t.Parallel()'d for the same
|
|
||||||
// reason — concurrent test goroutines would always show up.
|
|
||||||
func TestNoGoroutineLeaks(t *testing.T) {
|
func TestNoGoroutineLeaks(t *testing.T) {
|
||||||
defer goleak.VerifyNone(t, goleak.IgnoreCurrent())
|
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{})
|
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)
|
||||||
|
|||||||
@@ -0,0 +1,125 @@
|
|||||||
|
//go:build e2e_testing
|
||||||
|
// +build e2e_testing
|
||||||
|
|
||||||
|
package e2e
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/ed25519"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/pem"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/cert"
|
||||||
|
"github.com/slackhq/nebula/cert_test"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/ssh"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSSHDLifecycle(t *testing.T) {
|
||||||
|
// TestSSHDLifecycle exercises the in-process sshd through several config reloads and a Control.Stop.
|
||||||
|
ca, _, caKey, _ := cert_test.NewTestCaCert(
|
||||||
|
cert.Version1, cert.Curve_CURVE25519,
|
||||||
|
time.Now(), time.Now().Add(10*time.Minute),
|
||||||
|
nil, nil, []string{},
|
||||||
|
)
|
||||||
|
|
||||||
|
hostKeyPEM := generateSSHHostKey(t)
|
||||||
|
clientSigner, clientAuthKey := generateSSHClientKey(t)
|
||||||
|
sshdAddr := allocLoopbackPort(t)
|
||||||
|
|
||||||
|
overrides := m{
|
||||||
|
"sshd": m{
|
||||||
|
"enabled": true,
|
||||||
|
"listen": sshdAddr,
|
||||||
|
"host_key": hostKeyPEM,
|
||||||
|
"authorized_users": []m{{
|
||||||
|
"user": "tester",
|
||||||
|
"keys": []string{clientAuthKey},
|
||||||
|
}},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
control, _, _, _ := newSimpleServer(cert.Version1, ca, caKey, "sshd-test", "10.222.0.1/24", overrides)
|
||||||
|
control.Start()
|
||||||
|
t.Cleanup(func() { control.Stop() })
|
||||||
|
|
||||||
|
// sshd binds in a goroutine after Start returns; wait for it.
|
||||||
|
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"sshd never started listening")
|
||||||
|
|
||||||
|
for i := 1; i <= 3; i++ {
|
||||||
|
out := sshExecReload(t, sshdAddr, clientSigner)
|
||||||
|
assert.Contains(t, out, "Reloading config", "reload cycle %d", i)
|
||||||
|
require.Eventually(t, func() bool { return canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"sshd not listening after reload cycle %d", i)
|
||||||
|
}
|
||||||
|
|
||||||
|
control.Stop()
|
||||||
|
require.Eventually(t, func() bool { return !canDial(sshdAddr) }, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"sshd still listening after Control.Stop")
|
||||||
|
}
|
||||||
|
|
||||||
|
func canDial(addr string) bool {
|
||||||
|
c, err := net.DialTimeout("tcp", addr, 100*time.Millisecond)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = c.Close()
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// allocLoopbackPort grabs an unused TCP port on 127.0.0.1, closes it, and returns the address. There
|
||||||
|
// is a small race between releasing the port and the sshd reclaiming it; in practice the OS keeps the
|
||||||
|
// port available long enough for the test to bind it.
|
||||||
|
func allocLoopbackPort(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
l, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
require.NoError(t, err)
|
||||||
|
addr := l.Addr().String()
|
||||||
|
require.NoError(t, l.Close())
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateSSHHostKey(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
block, err := ssh.MarshalPrivateKey(priv, "nebula-e2e-host")
|
||||||
|
require.NoError(t, err)
|
||||||
|
return string(pem.EncodeToMemory(block))
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateSSHClientKey(t *testing.T) (ssh.Signer, string) {
|
||||||
|
t.Helper()
|
||||||
|
_, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
signer, err := ssh.NewSignerFromKey(priv)
|
||||||
|
require.NoError(t, err)
|
||||||
|
auth := strings.TrimSpace(string(ssh.MarshalAuthorizedKey(signer.PublicKey())))
|
||||||
|
return signer, auth
|
||||||
|
}
|
||||||
|
|
||||||
|
func sshExecReload(t *testing.T, addr string, signer ssh.Signer) string {
|
||||||
|
t.Helper()
|
||||||
|
cfg := &ssh.ClientConfig{
|
||||||
|
User: "tester",
|
||||||
|
Auth: []ssh.AuthMethod{ssh.PublicKeys(signer)},
|
||||||
|
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
||||||
|
Timeout: 2 * time.Second,
|
||||||
|
}
|
||||||
|
client, err := ssh.Dial("tcp", addr, cfg)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer client.Close()
|
||||||
|
|
||||||
|
sess, err := client.NewSession()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer sess.Close()
|
||||||
|
|
||||||
|
// reload tears the channel down before sending exit-status, so Output returns an error on the
|
||||||
|
// channel close. The output buffer still has whatever the reload callback wrote before that.
|
||||||
|
out, _ := sess.Output("reload")
|
||||||
|
return string(out)
|
||||||
|
}
|
||||||
@@ -138,6 +138,14 @@ 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.
|
||||||
@@ -163,17 +171,21 @@ 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
|
||||||
@@ -282,6 +294,24 @@ 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
|
||||||
|
|||||||
@@ -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.0
|
github.com/gaissmai/bart v0.26.1
|
||||||
github.com/gogo/protobuf v1.3.2
|
github.com/gogo/protobuf v1.3.2
|
||||||
github.com/google/gopacket v1.1.19
|
github.com/google/gopacket v1.1.19
|
||||||
github.com/kardianos/service v1.2.4
|
github.com/kardianos/service v1.2.4
|
||||||
@@ -26,7 +26,7 @@ require (
|
|||||||
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.52.0
|
golang.org/x/net v0.53.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
|
||||||
|
|||||||
@@ -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.0 h1:xOZ57E9hJLBiQaSyeZa9wgWhGuzfGACgqp4BE77OkO0=
|
github.com/gaissmai/bart v0.26.1 h1:+w4rnLGNlA2GDVn382Tfe3jOsK5vOr5n4KmigJ9lbTo=
|
||||||
github.com/gaissmai/bart v0.26.0/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
github.com/gaissmai/bart v0.26.1/go.mod h1:GREWQfTLRWz/c5FTOsIw+KkscuFkIV5t8Rp7Nd1Td5c=
|
||||||
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as=
|
||||||
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY=
|
||||||
@@ -182,8 +182,8 @@ golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLL
|
|||||||
golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA=
|
golang.org/x/net v0.0.0-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.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
|
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
|
||||||
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
|
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
|
||||||
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw=
|
||||||
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ type CertVerifier func(cert.Certificate) (*cert.CachedCertificate, error)
|
|||||||
type Result struct {
|
type Result struct {
|
||||||
EKey *noise.CipherState
|
EKey *noise.CipherState
|
||||||
DKey *noise.CipherState
|
DKey *noise.CipherState
|
||||||
|
Cipher noise.CipherFunc // identifies which post-handshake CipherState the data plane should wrap EKey/DKey in
|
||||||
MyCert cert.Certificate
|
MyCert cert.Certificate
|
||||||
RemoteCert *cert.CachedCertificate
|
RemoteCert *cert.CachedCertificate
|
||||||
RemoteIndex uint32
|
RemoteIndex uint32
|
||||||
@@ -105,6 +106,7 @@ func NewMachine(
|
|||||||
myVersion: version,
|
myVersion: version,
|
||||||
result: &Result{
|
result: &Result{
|
||||||
Initiator: initiator,
|
Initiator: initiator,
|
||||||
|
Cipher: cred.cipherSuite,
|
||||||
},
|
},
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -974,6 +974,7 @@ func (hm *HandshakeManager) continueHandshake(via ViaSender, hh *HandshakeHostIn
|
|||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
out := make([]byte, mtu)
|
out := make([]byte, mtu)
|
||||||
for _, cp := range hh.packetStore {
|
for _, cp := range hh.packetStore {
|
||||||
|
//todo use a sendbatcher
|
||||||
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
cp.callback(cp.messageType, cp.messageSubType, hostinfo, cp.packet, nb, out)
|
||||||
}
|
}
|
||||||
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
f.cachedPacketMetrics.sent.Inc(int64(len(hh.packetStore)))
|
||||||
|
|||||||
@@ -174,6 +174,10 @@ 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 {
|
||||||
@@ -185,6 +189,16 @@ 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,6 +2,7 @@ package nebula
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
@@ -9,10 +10,16 @@ import (
|
|||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
"github.com/slackhq/nebula/noiseutil"
|
"github.com/slackhq/nebula/noiseutil"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet, nb, out []byte, q int, localCache firewall.ConntrackCache) {
|
func (f *Interface) consumeInsidePacket(pkt wire.TunPacket, fwPacket *firewall.Packet, nb []byte, sendBatch *batch.SendBatch, rejectBuf []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
|
// borrowed: pkt.Bytes is owned by the originating tio.Queue and is
|
||||||
|
// only valid until the next Read on that queue. If you must keep
|
||||||
|
// the packet, use pkt.Clone() to detach it
|
||||||
|
packet := pkt.Bytes
|
||||||
err := newPacket(packet, false, fwPacket)
|
err := newPacket(packet, false, fwPacket)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
@@ -37,7 +44,10 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
// routes packets from the Nebula addr to the Nebula addr through the Nebula
|
||||||
// TUN device.
|
// TUN device.
|
||||||
if immediatelyForwardToSelf {
|
if immediatelyForwardToSelf {
|
||||||
_, err := f.readers[q].Write(packet)
|
err := pkt.PerSegment(func(seg []byte) error {
|
||||||
|
_, werr := f.readers[q].Write(seg)
|
||||||
|
return werr
|
||||||
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to forward to tun", "error", err)
|
f.l.Error("Failed to forward to tun", "error", err)
|
||||||
}
|
}
|
||||||
@@ -53,11 +63,23 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
|
|
||||||
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
hostinfo, ready := f.getOrHandshakeConsiderRouting(fwPacket, func(hh *HandshakeHostInfo) {
|
||||||
hh.cachePacket(f.l, header.Message, 0, packet, f.sendMessageNow, f.cachedPacketMetrics)
|
// borrowed: SegmentSuperpacket builds each segment in the kernel-supplied pkt
|
||||||
|
// bytes underneath. cachePacket explicitly copies its argument (handshake_manager.go cachePacket),
|
||||||
|
// so retaining segments past the loop is safe.
|
||||||
|
err := pkt.PerSegment(func(seg []byte) error {
|
||||||
|
hh.cachePacket(f.l, header.Message, 0, seg, f.sendMessageNow, f.cachedPacketMetrics)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil && f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Failed to segment superpacket for handshake cache",
|
||||||
|
"error", err,
|
||||||
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
|
)
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
if hostinfo == nil {
|
if hostinfo == nil {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
f.l.Debug("dropping outbound packet, vpnAddr not in our vpn networks or in unsafe networks",
|
||||||
"vpnAddr", fwPacket.RemoteAddr,
|
"vpnAddr", fwPacket.RemoteAddr,
|
||||||
@@ -73,10 +95,9 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
|
|
||||||
dropReason := f.firewall.Drop(*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.sendNoMetrics(header.Message, 0, hostinfo.ConnectionState, hostinfo, netip.AddrPort{}, packet, nb, out, q)
|
f.sendInsideMessage(hostinfo, pkt, nb, sendBatch)
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
f.rejectInside(packet, out, q)
|
f.rejectInside(packet, rejectBuf, q)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
hostinfo.logger(f.l).Debug("dropping outbound packet",
|
||||||
"fwPacket", fwPacket,
|
"fwPacket", fwPacket,
|
||||||
@@ -86,6 +107,124 @@ func (f *Interface) consumeInsidePacket(packet []byte, fwPacket *firewall.Packet
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (f *Interface) sendInsideEncrypt(hostinfo *HostInfo, ci *ConnectionState, seg, scratch, nb []byte) []byte {
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Lock()
|
||||||
|
}
|
||||||
|
c := ci.messageCounter.Add(1)
|
||||||
|
|
||||||
|
out := header.Encode(scratch, header.Version, header.Message, 0, hostinfo.remoteIndexId, c)
|
||||||
|
f.connectionManager.Out(hostinfo)
|
||||||
|
|
||||||
|
out, encErr := ci.eKey.EncryptDanger(out, out, seg, c, nb)
|
||||||
|
if noiseutil.EncryptLockNeeded {
|
||||||
|
ci.writeLock.Unlock()
|
||||||
|
}
|
||||||
|
if encErr != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to encrypt outgoing packet",
|
||||||
|
"error", encErr,
|
||||||
|
"udpAddr", hostinfo.remote,
|
||||||
|
"counter", c,
|
||||||
|
)
|
||||||
|
// Skip this segment; the rest of the superpacket can still
|
||||||
|
// go out — TCP will retransmit anything we drop here.
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendInsideMessage encrypts a firewall-approved inside packet (or every
|
||||||
|
// segment of a TSO/USO superpacket) into the caller's batch slot for
|
||||||
|
// later sendmmsg flush. Segmentation is fused with encryption here so the
|
||||||
|
// kernel-supplied superpacket bytes never get written into a separate
|
||||||
|
// scratch arena: PerSegment 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 wire.TunPacket, nb []byte, sendBatch *batch.SendBatch) {
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if ci.eKey == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if hostinfo.lastRebindCount != f.rebindCount {
|
||||||
|
//NOTE: there is an update hole if a tunnel isn't used and exactly 256 rebinds occur before the tunnel is
|
||||||
|
// finally used again. This tunnel would eventually be torn down and recreated if this action didn't help.
|
||||||
|
f.lightHouse.QueryServer(hostinfo.vpnAddrs[0])
|
||||||
|
hostinfo.lastRebindCount = f.rebindCount
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Lighthouse update triggered for punch due to rebind counter",
|
||||||
|
"vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !hostinfo.remote.IsValid() { //the relay path
|
||||||
|
//first, find our relay hostinfo:
|
||||||
|
var relayHostInfo *HostInfo
|
||||||
|
var relay *Relay
|
||||||
|
var err error
|
||||||
|
for _, relayIP := range hostinfo.relayState.CopyRelayIps() {
|
||||||
|
relayHostInfo, relay, err = f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relayIP)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.relayState.DeleteRelay(relayIP)
|
||||||
|
hostinfo.logger(f.l).Info("sendNoMetrics failed to find HostInfo",
|
||||||
|
"relay", relayIP,
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if relayHostInfo == nil || relay == nil {
|
||||||
|
//failure already logged
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err = pkt.PerSegment(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
|
||||||
|
}
|
||||||
|
|
||||||
|
sendBatch.Commit(toSend, relayHostInfo.remote, 0)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for relay send", "error", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
err := pkt.PerSegment(func(seg []byte) error {
|
||||||
|
// header + plaintext + AEAD tag (16 bytes for both AES-GCM and ChaCha20-Poly1305)
|
||||||
|
scratch := sendBatch.Reserve(header.Len + len(seg) + 16)
|
||||||
|
|
||||||
|
out := f.sendInsideEncrypt(hostinfo, ci, seg, scratch, nb)
|
||||||
|
if out == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sendBatch.Commit(out, hostinfo.remote, 0)
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Error("Failed to segment superpacket for send",
|
||||||
|
"error", err,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
func (f *Interface) rejectInside(packet []byte, out []byte, q int) {
|
||||||
if !f.firewall.InSendReject {
|
if !f.firewall.InSendReject {
|
||||||
return
|
return
|
||||||
@@ -275,21 +414,13 @@ func (f *Interface) sendTo(t header.MessageType, st header.MessageSubType, ci *C
|
|||||||
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
f.sendNoMetrics(t, st, ci, hostinfo, remote, p, nb, out, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
func (f *Interface) prepareSendVia(via *HostInfo,
|
||||||
// to the payload for the ultimate target host, making this a useful method for sending
|
|
||||||
// handshake messages to peers through relay tunnels.
|
|
||||||
// via is the HostInfo through which the message is relayed.
|
|
||||||
// ad is the plaintext data to authenticate, but not encrypt
|
|
||||||
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
|
||||||
// out is a buffer used to store the result of the Encrypt operation
|
|
||||||
// q indicates which writer to use to send the packet.
|
|
||||||
func (f *Interface) SendVia(via *HostInfo,
|
|
||||||
relay *Relay,
|
relay *Relay,
|
||||||
ad,
|
ad,
|
||||||
nb,
|
nb,
|
||||||
out []byte,
|
out []byte,
|
||||||
nocopy bool,
|
nocopy bool,
|
||||||
) {
|
) ([]byte, error) {
|
||||||
if noiseutil.EncryptLockNeeded {
|
if noiseutil.EncryptLockNeeded {
|
||||||
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
// NOTE: for goboring AESGCMTLS we need to lock because of the nonce check
|
||||||
via.ConnectionState.writeLock.Lock()
|
via.ConnectionState.writeLock.Lock()
|
||||||
@@ -311,7 +442,7 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
"headerLen", len(out),
|
"headerLen", len(out),
|
||||||
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
"cipherOverhead", via.ConnectionState.eKey.Overhead(),
|
||||||
)
|
)
|
||||||
return
|
return nil, io.ErrShortBuffer
|
||||||
}
|
}
|
||||||
|
|
||||||
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
// The header bytes are written to the 'out' slice; Grow the slice to hold the header and associated data payload.
|
||||||
@@ -331,13 +462,36 @@ func (f *Interface) SendVia(via *HostInfo,
|
|||||||
}
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
via.logger(f.l).Info("Failed to EncryptDanger in sendVia", "error", err)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
f.connectionManager.RelayUsed(relay.LocalIndex)
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendVia sends a payload through a Relay tunnel. No authentication or encryption is done
|
||||||
|
// to the payload for the ultimate target host, making this a useful method for sending
|
||||||
|
// handshake messages to peers through relay tunnels.
|
||||||
|
// via is the HostInfo through which the message is relayed.
|
||||||
|
// ad is the plaintext data to authenticate, but not encrypt
|
||||||
|
// nb is a buffer used to store the nonce value, re-used for performance reasons.
|
||||||
|
// out is a buffer used to store the result of the Encrypt operation
|
||||||
|
// q indicates which writer to use to send the packet.
|
||||||
|
func (f *Interface) SendVia(via *HostInfo,
|
||||||
|
relay *Relay,
|
||||||
|
ad,
|
||||||
|
nb,
|
||||||
|
out []byte,
|
||||||
|
nocopy bool,
|
||||||
|
) {
|
||||||
|
toSend, err := f.prepareSendVia(via, relay, ad, nb, out, nocopy)
|
||||||
|
if err != nil {
|
||||||
|
via.logger(f.l).Info("Failed to prepareSendVia", "error", err)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
err = f.writers[0].WriteTo(out, via.remote)
|
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) {
|
||||||
|
|||||||
+49
-19
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -13,11 +12,15 @@ import (
|
|||||||
|
|
||||||
"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/wire"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
"github.com/slackhq/nebula/firewall"
|
"github.com/slackhq/nebula/firewall"
|
||||||
"github.com/slackhq/nebula/header"
|
"github.com/slackhq/nebula/header"
|
||||||
"github.com/slackhq/nebula/overlay"
|
"github.com/slackhq/nebula/overlay"
|
||||||
|
"github.com/slackhq/nebula/overlay/batch"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -37,6 +40,7 @@ type InterfaceConfig struct {
|
|||||||
DropLocalBroadcast bool
|
DropLocalBroadcast bool
|
||||||
DropMulticast bool
|
DropMulticast bool
|
||||||
routines int
|
routines int
|
||||||
|
batchSize int
|
||||||
MessageMetrics *MessageMetrics
|
MessageMetrics *MessageMetrics
|
||||||
version string
|
version string
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
@@ -69,6 +73,7 @@ type Interface struct {
|
|||||||
dropLocalBroadcast bool
|
dropLocalBroadcast bool
|
||||||
dropMulticast bool
|
dropMulticast bool
|
||||||
routines int
|
routines int
|
||||||
|
batchSize int
|
||||||
disconnectInvalid atomic.Bool
|
disconnectInvalid atomic.Bool
|
||||||
closed atomic.Bool
|
closed atomic.Bool
|
||||||
relayManager *relayManager
|
relayManager *relayManager
|
||||||
@@ -88,8 +93,12 @@ type Interface struct {
|
|||||||
|
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
writers []udp.Conn
|
writers []udp.Conn
|
||||||
readers []io.ReadWriteCloser
|
readers []tio.Queue
|
||||||
wg sync.WaitGroup
|
// batchers is one per tun queue, wrapping readers[i].
|
||||||
|
// decryptToTun sends plaintext into the batch.RxBatcher;
|
||||||
|
// listenOut calls its Flush at the end of each UDP recvmmsg batch.
|
||||||
|
batchers []batch.RxBatcher
|
||||||
|
wg sync.WaitGroup
|
||||||
|
|
||||||
// fatalErr holds the first unexpected reader error that caused shutdown.
|
// fatalErr holds the first unexpected reader error that caused shutdown.
|
||||||
// nil means "no fatal error" (yet)
|
// nil means "no fatal error" (yet)
|
||||||
@@ -185,9 +194,11 @@ func NewInterface(ctx context.Context, c *InterfaceConfig) (*Interface, error) {
|
|||||||
dropLocalBroadcast: c.DropLocalBroadcast,
|
dropLocalBroadcast: c.DropLocalBroadcast,
|
||||||
dropMulticast: c.DropMulticast,
|
dropMulticast: c.DropMulticast,
|
||||||
routines: c.routines,
|
routines: c.routines,
|
||||||
|
batchSize: c.batchSize,
|
||||||
version: c.version,
|
version: c.version,
|
||||||
writers: make([]udp.Conn, c.routines),
|
writers: make([]udp.Conn, c.routines),
|
||||||
readers: make([]io.ReadWriteCloser, c.routines),
|
readers: make([]tio.Queue, c.routines),
|
||||||
|
batchers: make([]batch.RxBatcher, c.routines),
|
||||||
myVpnNetworks: cs.myVpnNetworks,
|
myVpnNetworks: cs.myVpnNetworks,
|
||||||
myVpnNetworksTable: cs.myVpnNetworksTable,
|
myVpnNetworksTable: cs.myVpnNetworksTable,
|
||||||
myVpnAddrs: cs.myVpnAddrs,
|
myVpnAddrs: cs.myVpnAddrs,
|
||||||
@@ -245,15 +256,17 @@ func (f *Interface) activate() error {
|
|||||||
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
metrics.GetOrRegisterGauge("routines", nil).Update(int64(f.routines))
|
||||||
|
|
||||||
// Prepare n tun queues
|
// Prepare n tun queues
|
||||||
var reader io.ReadWriteCloser = f.inside
|
|
||||||
for i := 0; i < f.routines; i++ {
|
for i := 0; i < f.routines; i++ {
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
reader, err = f.inside.NewMultiQueueReader()
|
if err = f.inside.NewMultiQueueReader(); err != nil {
|
||||||
if err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
f.readers[i] = reader
|
}
|
||||||
|
f.readers = f.inside.Readers()
|
||||||
|
for i := range f.readers {
|
||||||
|
arena := util.NewArena(max(f.batchSize, 1) * udp.MTU)
|
||||||
|
f.batchers[i] = batch.NewPassthrough(f.readers[i], f.batchSize, arena)
|
||||||
}
|
}
|
||||||
|
|
||||||
f.wg.Add(1) // for us to wait on Close() to return
|
f.wg.Add(1) // for us to wait on Close() to return
|
||||||
@@ -311,14 +324,21 @@ func (f *Interface) listenOut(i int) {
|
|||||||
|
|
||||||
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
ctCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
lhh := f.lightHouse.NewRequestHandler()
|
lhh := f.lightHouse.NewRequestHandler()
|
||||||
plaintext := make([]byte, udp.MTU)
|
|
||||||
h := &header.H{}
|
h := &header.H{}
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
err := li.ListenOut(func(fromUdpAddr netip.AddrPort, payload []byte) {
|
listener := func(fromUdpAddr netip.AddrPort, payload []byte, meta udp.RxMeta) {
|
||||||
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, plaintext[:0], payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
f.readOutsidePackets(ViaSender{UdpAddr: fromUdpAddr}, payload, h, fwPacket, lhh, nb, i, ctCache.Get())
|
||||||
})
|
}
|
||||||
|
|
||||||
|
flusher := func() {
|
||||||
|
if err := f.batchers[i].Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to flush tun coalescer", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
err := li.ListenOut(listener, flusher)
|
||||||
|
|
||||||
if err != nil && !f.closed.Load() {
|
if err != nil && !f.closed.Load() {
|
||||||
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
f.l.Error("Error while reading inbound packet, closing", "error", err)
|
||||||
@@ -328,28 +348,38 @@ func (f *Interface) listenOut(i int) {
|
|||||||
f.l.Debug("underlay reader is done", "reader", i)
|
f.l.Debug("underlay reader is done", "reader", i)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) listenIn(reader io.ReadWriteCloser, i int) {
|
func (f *Interface) listenIn(reader tio.Queue, q int) {
|
||||||
packet := make([]byte, mtu)
|
packetMem := make([]byte, mtu+16) //MTU + some leading slack space for platforms that return "bonus info"
|
||||||
out := make([]byte, mtu)
|
// TODO get the amount of bonus info from the reader
|
||||||
|
packets := make([]wire.TunPacket, 1)
|
||||||
|
rejectBuf := make([]byte, mtu)
|
||||||
|
arenaSize := batch.SendBatchCap * (udp.MTU + 32)
|
||||||
|
sb := batch.NewSendBatch(f.writers[q], batch.SendBatchCap, util.NewArena(arenaSize))
|
||||||
fwPacket := &firewall.Packet{}
|
fwPacket := &firewall.Packet{}
|
||||||
nb := make([]byte, 12, 12)
|
nb := make([]byte, 12, 12)
|
||||||
|
|
||||||
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
conntrackCache := firewall.NewConntrackCacheTicker(f.ctx, f.l, f.conntrackCacheTimeout)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
n, err := reader.Read(packet)
|
n, err := reader.Read(packets, packetMem)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !f.closed.Load() {
|
if !f.closed.Load() {
|
||||||
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", i)
|
f.l.Error("Error while reading outbound packet, closing", "error", err, "reader", q)
|
||||||
f.onFatal(err)
|
f.onFatal(err)
|
||||||
}
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|
||||||
f.consumeInsidePacket(packet[:n], fwPacket, nb, out, i, conntrackCache.Get())
|
ctCache := conntrackCache.Get()
|
||||||
|
for i := range n {
|
||||||
|
f.consumeInsidePacket(packets[i], fwPacket, nb, sb, rejectBuf, q, ctCache)
|
||||||
|
}
|
||||||
|
if err := sb.Flush(); err != nil {
|
||||||
|
f.l.Error("Failed to write outgoing batch", "error", err, "writer", q)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
f.l.Debug("overlay reader is done", "reader", i)
|
f.l.Debug("overlay reader is done", "reader", q)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
func (f *Interface) RegisterConfigChangeCallbacks(c *config.C) {
|
||||||
|
|||||||
+6
-46
@@ -15,7 +15,6 @@ 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"
|
||||||
@@ -35,7 +34,6 @@ 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
|
||||||
@@ -75,9 +73,8 @@ 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
|
||||||
metricHolepunchTx metrics.Counter
|
l *slog.Logger
|
||||||
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
|
||||||
@@ -105,7 +102,6 @@ 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)),
|
||||||
@@ -118,9 +114,6 @@ 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)
|
||||||
@@ -1406,58 +1399,25 @@ 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()) {
|
||||||
punch(b, detailsVpnAddr)
|
lhh.lh.punchy.Schedule(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()) {
|
||||||
punch(b, detailsVpnAddr)
|
lhh.lh.punchy.Schedule(b, detailsVpnAddr)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// This sends a nebula test packet to the host trying to contact us. In the case
|
// 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.
|
// a tunnel. ScheduleRespond is a no-op when punchy.respond is disabled.
|
||||||
if lhh.lh.punchy.GetRespond() {
|
lhh.lh.punchy.ScheduleRespond(detailsVpnAddr)
|
||||||
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 {
|
||||||
|
|||||||
@@ -55,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(l.With("subsystem", "sshd"))
|
ssh, err := sshd.NewSSHServer(ctx, 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)
|
||||||
}
|
}
|
||||||
@@ -170,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)
|
punchy := NewPunchyFromConfig(l, c, udpConns[0])
|
||||||
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 {
|
||||||
@@ -215,6 +215,7 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
DropLocalBroadcast: c.GetBool("tun.drop_local_broadcast", false),
|
DropLocalBroadcast: c.GetBool("tun.drop_local_broadcast", false),
|
||||||
DropMulticast: c.GetBool("tun.drop_multicast", false),
|
DropMulticast: c.GetBool("tun.drop_multicast", false),
|
||||||
routines: routines,
|
routines: routines,
|
||||||
|
batchSize: c.GetInt("listen.batch", 64),
|
||||||
MessageMetrics: messageMetrics,
|
MessageMetrics: messageMetrics,
|
||||||
version: buildVersion,
|
version: buildVersion,
|
||||||
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
relayManager: NewRelayManager(ctx, l, hostMap, c),
|
||||||
@@ -240,6 +241,8 @@ func Main(c *config.C, configTest bool, buildVersion string, l *slog.Logger, dev
|
|||||||
|
|
||||||
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)
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ 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) {
|
||||||
@@ -33,6 +35,11 @@ func (m *MessageMetrics) Tx(t header.MessageType, s header.MessageSubType, i int
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
func (m *MessageMetrics) RxInvalid(i int64) {
|
||||||
|
if m != nil && m.rxInvalid != nil {
|
||||||
|
m.rxInvalid.Inc(i)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func newMessageMetrics() *MessageMetrics {
|
func newMessageMetrics() *MessageMetrics {
|
||||||
gen := func(t string) [][]metrics.Counter {
|
gen := func(t string) [][]metrics.Counter {
|
||||||
@@ -56,6 +63,7 @@ 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),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,73 +0,0 @@
|
|||||||
package nebula
|
|
||||||
|
|
||||||
import (
|
|
||||||
"crypto/cipher"
|
|
||||||
"encoding/binary"
|
|
||||||
"errors"
|
|
||||||
|
|
||||||
"github.com/flynn/noise"
|
|
||||||
)
|
|
||||||
|
|
||||||
type endianness interface {
|
|
||||||
PutUint64(b []byte, v uint64)
|
|
||||||
}
|
|
||||||
|
|
||||||
var noiseEndianness endianness = binary.BigEndian
|
|
||||||
|
|
||||||
type NebulaCipherState struct {
|
|
||||||
c 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
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
package noiseutil
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CipherStateAESGCM is the data-plane wrapper for the AES-GCM AEAD cipher.
|
||||||
|
// AES-GCM uses big-endian nonce encoding per the Noise spec.
|
||||||
|
type CipherStateAESGCM struct {
|
||||||
|
c cipher.AEAD
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCipherStateAESGCM extracts the underlying AEAD from the post-handshake noise.CipherState.
|
||||||
|
// The caller is responsible for ensuring the noise cipher is actually AES-GCM,
|
||||||
|
// otherwise the type assertion still succeeds but the nonce endianness will be wrong on the wire.
|
||||||
|
func NewCipherStateAESGCM(s *noise.CipherState) *CipherStateAESGCM {
|
||||||
|
return &CipherStateAESGCM{c: s.Cipher().(cipher.AEAD)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *CipherStateAESGCM) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||||
|
if s == nil {
|
||||||
|
return nil, errors.New("no cipher state available to encrypt")
|
||||||
|
}
|
||||||
|
nb[0] = 0
|
||||||
|
nb[1] = 0
|
||||||
|
nb[2] = 0
|
||||||
|
nb[3] = 0
|
||||||
|
binary.BigEndian.PutUint64(nb[4:], n)
|
||||||
|
return s.c.Seal(out, nb, plaintext, ad), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *CipherStateAESGCM) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||||
|
if s == nil {
|
||||||
|
return []byte{}, nil
|
||||||
|
}
|
||||||
|
nb[0] = 0
|
||||||
|
nb[1] = 0
|
||||||
|
nb[2] = 0
|
||||||
|
nb[3] = 0
|
||||||
|
binary.BigEndian.PutUint64(nb[4:], n)
|
||||||
|
return s.c.Open(out, nb, ciphertext, ad)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *CipherStateAESGCM) Overhead() int {
|
||||||
|
if s == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return s.c.Overhead()
|
||||||
|
}
|
||||||
@@ -0,0 +1,52 @@
|
|||||||
|
package noiseutil
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CipherStateChaChaPoly is the data-plane wrapper for the ChaCha20-Poly1305 AEAD cipher.
|
||||||
|
// ChaCha20-Poly1305 uses little-endian nonce encoding per the Noise spec.
|
||||||
|
type CipherStateChaChaPoly struct {
|
||||||
|
c cipher.AEAD
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCipherStateChaChaPoly extracts the underlying AEAD from the post-handshake noise.CipherState.
|
||||||
|
// The caller is responsible for ensuring the noise cipher is actually ChaCha20-Poly1305.
|
||||||
|
func NewCipherStateChaChaPoly(s *noise.CipherState) *CipherStateChaChaPoly {
|
||||||
|
return &CipherStateChaChaPoly{c: s.Cipher().(cipher.AEAD)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *CipherStateChaChaPoly) EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||||
|
if s == nil {
|
||||||
|
return nil, errors.New("no cipher state available to encrypt")
|
||||||
|
}
|
||||||
|
nb[0] = 0
|
||||||
|
nb[1] = 0
|
||||||
|
nb[2] = 0
|
||||||
|
nb[3] = 0
|
||||||
|
binary.LittleEndian.PutUint64(nb[4:], n)
|
||||||
|
return s.c.Seal(out, nb, plaintext, ad), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *CipherStateChaChaPoly) DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error) {
|
||||||
|
if s == nil {
|
||||||
|
return []byte{}, nil
|
||||||
|
}
|
||||||
|
nb[0] = 0
|
||||||
|
nb[1] = 0
|
||||||
|
nb[2] = 0
|
||||||
|
nb[3] = 0
|
||||||
|
binary.LittleEndian.PutUint64(nb[4:], n)
|
||||||
|
return s.c.Open(out, nb, ciphertext, ad)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *CipherStateChaChaPoly) Overhead() int {
|
||||||
|
if s == nil {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return s.c.Overhead()
|
||||||
|
}
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
package noiseutil
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/flynn/noise"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CipherState is the post-handshake AEAD cipher used for the data plane.
|
||||||
|
// Each supported cipher has its own concrete implementation in this package with the nonce endianness hardcoded,
|
||||||
|
// so the encrypt/decrypt fast path avoids interface dispatch on the byte order.
|
||||||
|
type CipherState interface {
|
||||||
|
// EncryptDanger encrypts and authenticates a given payload.
|
||||||
|
//
|
||||||
|
// out is a destination slice to hold the output of the EncryptDanger operation.
|
||||||
|
// - ad is additional data, which will be authenticated and appended to out, but not encrypted.
|
||||||
|
// - plaintext is encrypted, authenticated and appended to out.
|
||||||
|
// - n is a nonce value which must never be re-used with this key.
|
||||||
|
// - nb is a scratch buffer used to assemble the nonce.
|
||||||
|
EncryptDanger(out, ad, plaintext []byte, n uint64, nb []byte) ([]byte, error)
|
||||||
|
|
||||||
|
// DecryptDanger authenticates and decrypts a given payload, with the same argument shape as EncryptDanger.
|
||||||
|
DecryptDanger(out, ad, ciphertext []byte, n uint64, nb []byte) ([]byte, error)
|
||||||
|
|
||||||
|
// Overhead returns the AEAD tag size, or 0 if the receiver is nil.
|
||||||
|
Overhead() int
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewCipherState wraps the post-handshake noise.CipherState in the per-cipher type that matches cipherFunc.
|
||||||
|
// cipherFunc must be the same cipher used to build the noise CipherSuite that produced s.
|
||||||
|
func NewCipherState(s *noise.CipherState, cipherFunc noise.CipherFunc) CipherState {
|
||||||
|
switch cipherFunc.CipherName() {
|
||||||
|
case CipherAESGCM.CipherName():
|
||||||
|
return NewCipherStateAESGCM(s)
|
||||||
|
case noise.CipherChaChaPoly.CipherName():
|
||||||
|
return NewCipherStateChaChaPoly(s)
|
||||||
|
default:
|
||||||
|
panic(fmt.Sprintf("noiseutil: unsupported cipher %q", cipherFunc.CipherName()))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,222 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestDecryptDangerRelayShapeNoAlloc covers the AD-only relay path used in
|
||||||
|
// outside.go's handleOutsideRelayPacket: the body is AD, the trailing 16 bytes
|
||||||
|
// are the AEAD tag, the plaintext is empty, and the caller passes nil as the
|
||||||
|
// destination because it only needs the auth side-effect. The call must
|
||||||
|
// succeed, return an empty plaintext, and not allocate on the hot path.
|
||||||
|
func TestDecryptDangerRelayShapeNoAlloc(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
c noise.CipherFunc
|
||||||
|
wrap func(*noise.CipherState) CipherState
|
||||||
|
}{
|
||||||
|
{"AESGCM", CipherAESGCM, func(cs *noise.CipherState) CipherState { return NewCipherStateAESGCM(cs) }},
|
||||||
|
{"ChaChaPoly", noise.CipherChaChaPoly, func(cs *noise.CipherState) CipherState { return NewCipherStateChaChaPoly(cs) }},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
encCS, decCS := buildCipherStates(t, tc.c)
|
||||||
|
enc, dec := tc.wrap(encCS), tc.wrap(decCS)
|
||||||
|
|
||||||
|
ad := make([]byte, 1200) // typical relay packet body size
|
||||||
|
for i := range ad {
|
||||||
|
ad[i] = byte(i)
|
||||||
|
}
|
||||||
|
nb := make([]byte, 12)
|
||||||
|
|
||||||
|
// Build the "signature value" the way handleOutsideRelayPacket sees it:
|
||||||
|
// empty plaintext encrypted with the body as AD yields just the 16-byte tag.
|
||||||
|
tag, err := enc.EncryptDanger(nil, ad, nil, 1, nb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, tag, dec.Overhead())
|
||||||
|
|
||||||
|
// Sanity: the relay-shaped call returns empty plaintext, no error.
|
||||||
|
out, err := dec.DecryptDanger(nil, ad, tag, 1, nb)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Empty(t, out)
|
||||||
|
|
||||||
|
// Tampering with the AD must fail authentication.
|
||||||
|
ad[0] ^= 0xff
|
||||||
|
_, err = dec.DecryptDanger(nil, ad, tag, 1, nb)
|
||||||
|
require.Error(t, err)
|
||||||
|
ad[0] ^= 0xff
|
||||||
|
|
||||||
|
// The hot path must not allocate. AllocsPerRun does a warm-up run, so any
|
||||||
|
// one-time setup is excluded. Counter has to advance so the AEAD nonce is
|
||||||
|
// unique per call, but we don't care whether the auth succeeds — we only
|
||||||
|
// care about whether the call path allocates.
|
||||||
|
var counter uint64 = 2
|
||||||
|
allocs := testing.AllocsPerRun(100, func() {
|
||||||
|
_, _ = dec.DecryptDanger(nil, ad, tag, counter, nb)
|
||||||
|
counter++
|
||||||
|
})
|
||||||
|
assert.Equal(t, 0.0, allocs, "DecryptDanger(nil, ...) must not allocate")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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())
|
||||||
|
}
|
||||||
+192
-227
@@ -20,23 +20,46 @@ const (
|
|||||||
minFwPacketLen = 4
|
minFwPacketLen = 4
|
||||||
)
|
)
|
||||||
|
|
||||||
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) {
|
var ErrOutOfWindow = errors.New("out of window packet")
|
||||||
|
|
||||||
|
func (f *Interface) readOutsidePackets(via ViaSender, 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.l.Info("Error while parsing inbound packet",
|
f.messageMetrics.RxInvalid(1)
|
||||||
"from", via,
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
"error", err,
|
f.l.Debug("Error while parsing inbound packet",
|
||||||
"packet", packet,
|
"from", via,
|
||||||
)
|
"error", err,
|
||||||
|
"packet", packet,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if h.Version != header.Version {
|
||||||
|
f.messageMetrics.RxInvalid(1)
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Unexpected header version received", "from", via)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check before processing to see if this is a expected type/subtype
|
||||||
|
if !h.IsValidSubType() {
|
||||||
|
f.messageMetrics.RxInvalid(1)
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
f.l.Debug("Unexpected packet received", "from", via)
|
||||||
}
|
}
|
||||||
return
|
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)
|
||||||
}
|
}
|
||||||
@@ -44,215 +67,194 @@ 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
|
||||||
// verify if we've seen this index before, otherwise respond to the handshake initiation
|
if isMessageRelay {
|
||||||
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)
|
||||||
}
|
}
|
||||||
|
|
||||||
var ci *ConnectionState
|
// At this point we should have a valid existing tunnel, verify and send
|
||||||
if hostinfo != nil {
|
// recvError if necessary
|
||||||
ci = hostinfo.ConnectionState
|
if hostinfo == nil || hostinfo.ConnectionState == nil {
|
||||||
|
if !via.IsRelayed {
|
||||||
|
f.maybeSendRecvError(via.UdpAddr, h.RemoteIndex)
|
||||||
|
}
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// All remaining packets are encrypted
|
||||||
|
ci := hostinfo.ConnectionState
|
||||||
|
if !ci.window.Check(f.l, h.MessageCounter) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Relay packets are special
|
||||||
|
if isMessageRelay {
|
||||||
|
f.handleOutsideRelayPacket(hostinfo, via, packet, h, fwPacket, lhf, nb, q, localCache)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
out := f.batchers[q].Reserve(len(packet))[:0]
|
||||||
|
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:
|
||||||
if !f.decryptToTun(hostinfo, h.MessageCounter, out, packet, fwPacket, nb, q, localCache) {
|
f.handleOutsideMessagePacket(hostinfo, out, packet, fwPacket, nb, q, localCache)
|
||||||
return
|
default:
|
||||||
}
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message subtype seen", "from", via, "header", h)
|
||||||
case header.MessageRelay:
|
return
|
||||||
// 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, d, f)
|
lhf.HandleRequest(via.UdpAddr, hostinfo.vpnAddrs, out, f)
|
||||||
|
|
||||||
// Fallthrough to the bottom to record incoming traffic
|
|
||||||
|
|
||||||
case header.Test:
|
case header.Test:
|
||||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
switch h.Subtype {
|
||||||
if !f.handleEncrypted(ci, via, h) {
|
case header.TestReply:
|
||||||
|
// No-op, useful for the Roaming and connectionManager side-effects above
|
||||||
|
case header.TestRequest:
|
||||||
|
f.send(header.Test, header.TestReply, ci, hostinfo, out, nb, out)
|
||||||
|
default:
|
||||||
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected test subtype seen", "from", via, "header", h)
|
||||||
return
|
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:
|
||||||
if !f.handleEncrypted(ci, via, h) {
|
f.relayManager.HandleControlMsg(hostinfo, out, f)
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
d, err := f.decrypt(hostinfo, h.MessageCounter, out, packet, h, nb)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt Control packet",
|
|
||||||
"error", err,
|
|
||||||
"from", via,
|
|
||||||
"packet", packet,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f.relayManager.HandleControlMsg(hostinfo, d, f)
|
|
||||||
|
|
||||||
default:
|
default:
|
||||||
f.messageMetrics.Rx(h.Type, h.Subtype, 1)
|
hostinfo.logger(f.l).Error("IsValidSubType was true, but unexpected message type seen", "from", via, "header", h)
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
}
|
||||||
hostinfo.logger(f.l).Debug("Unexpected packet received", "from", via)
|
}
|
||||||
}
|
|
||||||
|
func (f *Interface) handleOutsideRelayPacket(hostinfo *HostInfo, via ViaSender, packet []byte, h *header.H, fwPacket *firewall.Packet, lhf *LightHouseHandler, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
|
// 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():]
|
||||||
|
// The decrypted output is empty (relay packets carry their payload as AD) and unused.
|
||||||
|
// The recursive readOutsidePackets call below operates on signedPayload. Passing
|
||||||
|
// nil avoids reserving an arena slot.
|
||||||
|
if _, err := hostinfo.ConnectionState.dKey.DecryptDanger(nil, signedPayload, signatureValue, h.MessageCounter, nb); 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
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.handleHostRoaming(hostinfo, via)
|
switch relay.Type {
|
||||||
|
case TerminalType:
|
||||||
|
// If I am the target of this relay, process the unwrapped packet
|
||||||
|
// From this recursive point, all these variables are 'burned'. We shouldn't rely on them again.
|
||||||
|
via = ViaSender{
|
||||||
|
UdpAddr: via.UdpAddr,
|
||||||
|
relayHI: hostinfo,
|
||||||
|
remoteIdx: relay.RemoteIndex,
|
||||||
|
relay: relay,
|
||||||
|
IsRelayed: true,
|
||||||
|
}
|
||||||
|
f.readOutsidePackets(via, signedPayload, h, fwPacket, lhf, nb, q, localCache)
|
||||||
|
case ForwardingType:
|
||||||
|
// Find the target HostInfo relay object
|
||||||
|
targetHI, targetRelay, err := f.hostMap.QueryVpnAddrsRelayFor(hostinfo.vpnAddrs, relay.PeerAddr)
|
||||||
|
if err != nil {
|
||||||
|
hostinfo.logger(f.l).Info("Failed to find target host info by ip",
|
||||||
|
"relayTo", relay.PeerAddr,
|
||||||
|
"error", err,
|
||||||
|
"hostinfo.vpnAddrs", hostinfo.vpnAddrs,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
f.connectionManager.In(hostinfo)
|
// If that relay is Established, forward the payload through it
|
||||||
|
if targetRelay.State == Established {
|
||||||
|
switch targetRelay.Type {
|
||||||
|
case ForwardingType:
|
||||||
|
// Forward this packet through the relay tunnel
|
||||||
|
// Find the target HostInfo //todo it would potentially be nice to batch these
|
||||||
|
out := f.batchers[q].Reserve(len(packet) + header.Len + hostinfo.ConnectionState.dKey.Overhead())[:0]
|
||||||
|
f.SendVia(targetHI, targetRelay, signedPayload, nb, out, false)
|
||||||
|
case TerminalType:
|
||||||
|
hostinfo.logger(f.l).Error("Unexpected Relay Type of Terminal")
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Unexpected targetRelay Type", "from", via, "relayType", targetRelay.Type)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
hostinfo.logger(f.l).Info("Unexpected target relay state",
|
||||||
|
"relayTo", relay.PeerAddr,
|
||||||
|
"relayFrom", hostinfo.vpnAddrs[0],
|
||||||
|
"targetRelayState", targetRelay.State,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
hostinfo.logger(f.l).Debug("Unexpected relay type", "from", via, "relayType", relay.Type)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
// closeTunnel closes a tunnel locally, it does not send a closeTunnel packet to the remote
|
||||||
@@ -300,23 +302,6 @@ func (f *Interface) handleHostRoaming(hostinfo *HostInfo, via ViaSender) {
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// handleEncrypted returns true if a packet should be processed, false otherwise
|
|
||||||
func (f *Interface) handleEncrypted(ci *ConnectionState, via ViaSender, h *header.H) bool {
|
|
||||||
// If connectionstate does not exist, send a recv error, if possible, to encourage a fast reconnect
|
|
||||||
if ci == nil {
|
|
||||||
if !via.IsRelayed {
|
|
||||||
f.maybeSendRecvError(via.UdpAddr, h.RemoteIndex)
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
// If the window check fails, refuse to process the packet, but don't send a recv error
|
|
||||||
if !ci.window.Check(f.l, h.MessageCounter) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
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")
|
||||||
@@ -523,38 +508,20 @@ 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) {
|
||||||
if f.l.Enabled(context.Background(), slog.LevelDebug) {
|
return nil, ErrOutOfWindow
|
||||||
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
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) bool {
|
func (f *Interface) handleOutsideMessagePacket(hostinfo *HostInfo, out []byte, packet []byte, fwPacket *firewall.Packet, nb []byte, q int, localCache firewall.ConntrackCache) {
|
||||||
var err error
|
err := newPacket(out, true, fwPacket)
|
||||||
|
|
||||||
out, err = hostinfo.ConnectionState.dKey.DecryptDanger(out, packet[:header.Len], packet[header.Len:], messageCounter, nb)
|
|
||||||
if err != nil {
|
|
||||||
hostinfo.logger(f.l).Error("Failed to decrypt packet", "error", err)
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
err = newPacket(out, true, fwPacket)
|
|
||||||
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 false
|
return
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
dropReason := f.firewall.Drop(*fwPacket, true, hostinfo, f.pki.GetCAPool(), localCache)
|
||||||
@@ -568,15 +535,13 @@ func (f *Interface) decryptToTun(hostinfo *HostInfo, messageCounter uint64, out
|
|||||||
"reason", dropReason,
|
"reason", dropReason,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
return false
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
f.connectionManager.In(hostinfo)
|
err = f.batchers[q].Commit(out)
|
||||||
_, err = f.readers[q].Write(out)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
f.l.Error("Failed to write to tun", "error", err)
|
f.l.Error("Failed to write to tun", "error", err)
|
||||||
}
|
}
|
||||||
return true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
func (f *Interface) maybeSendRecvError(endpoint netip.AddrPort, index uint32) {
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
// Arena is an injectable byte-slab that hands out non-overlapping borrowed
|
||||||
|
// slices via Reserve and releases them in bulk via Reset. Coalescers take
|
||||||
|
// an *Arena at construction so the caller controls the slab lifetime and
|
||||||
|
// can share one slab across multiple coalescers (MultiCoalescer hands the
|
||||||
|
// same *Arena to every lane so the lanes don't carry their own backings).
|
||||||
|
//
|
||||||
|
// Reserve borrows; the slice is valid until the next Reset. The slab grows
|
||||||
|
// (by allocating a fresh, larger backing array) if a Reserve doesn't fit;
|
||||||
|
// pre-size the arena via NewArena to avoid that path on the hot path.
|
||||||
|
type Arena struct {
|
||||||
|
buf []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewArena returns an Arena with a pre-allocated backing of the given
|
||||||
|
// capacity. Pass 0 if you don't intend to call Reserve (e.g. a test that
|
||||||
|
// only feeds the coalescer pre-made []byte packets via Commit).
|
||||||
|
func NewArena(capacity int) *Arena {
|
||||||
|
return &Arena{buf: make([]byte, 0, capacity)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reserve hands out a non-overlapping sz-byte slice from the arena. If the
|
||||||
|
// request doesn't fit the current backing, a fresh, larger backing is
|
||||||
|
// allocated; already-borrowed slices reference the old backing and remain
|
||||||
|
// valid until Reset.
|
||||||
|
func (a *Arena) Reserve(sz int) []byte {
|
||||||
|
if len(a.buf)+sz > cap(a.buf) {
|
||||||
|
newCap := max(cap(a.buf)*2, sz)
|
||||||
|
a.buf = make([]byte, 0, newCap)
|
||||||
|
}
|
||||||
|
start := len(a.buf)
|
||||||
|
a.buf = a.buf[:start+sz]
|
||||||
|
return a.buf[start : start+sz : start+sz]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reset releases every slice handed out since the last Reset. Callers must
|
||||||
|
// not use any previously-borrowed slice after this returns. The underlying
|
||||||
|
// backing array is retained so subsequent Reserves don't re-allocate.
|
||||||
|
func (a *Arena) Reset() {
|
||||||
|
a.buf = a.buf[:0]
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Passthrough is a RxBatcher that doesn't batch anything, it just accumulates and then sends packets.
|
||||||
|
type Passthrough struct {
|
||||||
|
out io.Writer
|
||||||
|
slots [][]byte
|
||||||
|
arena *util.Arena
|
||||||
|
cursor int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewPassthrough(w io.Writer, slots int, arena *util.Arena) *Passthrough {
|
||||||
|
return &Passthrough{
|
||||||
|
out: w,
|
||||||
|
slots: make([][]byte, 0, slots),
|
||||||
|
arena: arena,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Reserve(sz int) []byte {
|
||||||
|
return p.arena.Reserve(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Commit(pkt []byte) error {
|
||||||
|
p.slots = append(p.slots, pkt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Passthrough) Flush() error {
|
||||||
|
var firstErr error
|
||||||
|
for _, s := range p.slots {
|
||||||
|
_, err := p.out.Write(s)
|
||||||
|
if err != nil && firstErr == nil {
|
||||||
|
firstErr = err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
clear(p.slots)
|
||||||
|
p.slots = p.slots[:0]
|
||||||
|
p.arena.Reset()
|
||||||
|
return firstErr
|
||||||
|
}
|
||||||
@@ -0,0 +1,12 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
type RxBatcher interface {
|
||||||
|
// Reserve creates a pkt to borrow
|
||||||
|
Reserve(sz int) []byte
|
||||||
|
// Commit borrows pkt. The caller must keep pkt valid until the next Flush
|
||||||
|
Commit(pkt []byte) error
|
||||||
|
// Flush emits every queued packet in arrival order. Returns the
|
||||||
|
// first error observed; keeps draining so one bad packet doesn't hold up
|
||||||
|
// the rest. After Flush returns, borrowed payload slices may be recycled.
|
||||||
|
Flush() error
|
||||||
|
}
|
||||||
@@ -0,0 +1,60 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
const SendBatchCap = 128
|
||||||
|
|
||||||
|
// batchWriter is the minimal subset of udp.Conn needed by SendBatch to flush.
|
||||||
|
type batchWriter interface {
|
||||||
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SendBatch accumulates encrypted UDP packets and flushes them via WriteBatch.
|
||||||
|
// One SendBatch is owned by each listenIn goroutine; no locking is needed.
|
||||||
|
// Slot bytes are borrowed from the injected Arena and remain valid until
|
||||||
|
// Flush, which Resets the arena.
|
||||||
|
type SendBatch struct {
|
||||||
|
out batchWriter
|
||||||
|
bufs [][]byte
|
||||||
|
dsts []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
arena *util.Arena
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewSendBatch makes a SendBatch with batchCap slots backed by arena.
|
||||||
|
func NewSendBatch(out batchWriter, batchCap int, arena *util.Arena) *SendBatch {
|
||||||
|
return &SendBatch{
|
||||||
|
out: out,
|
||||||
|
bufs: make([][]byte, 0, batchCap),
|
||||||
|
dsts: make([]netip.AddrPort, 0, batchCap),
|
||||||
|
ecns: make([]byte, 0, batchCap),
|
||||||
|
arena: arena,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Reserve(sz int) []byte {
|
||||||
|
return b.arena.Reserve(sz)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Commit(pkt []byte, dst netip.AddrPort, outerECN byte) {
|
||||||
|
b.bufs = append(b.bufs, pkt)
|
||||||
|
b.dsts = append(b.dsts, dst)
|
||||||
|
b.ecns = append(b.ecns, outerECN)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *SendBatch) Flush() error {
|
||||||
|
var err error
|
||||||
|
if len(b.bufs) > 0 {
|
||||||
|
err = b.out.WriteBatch(b.bufs, b.dsts, b.ecns)
|
||||||
|
}
|
||||||
|
clear(b.bufs)
|
||||||
|
b.bufs = b.bufs[:0]
|
||||||
|
b.dsts = b.dsts[:0]
|
||||||
|
b.ecns = b.ecns[:0]
|
||||||
|
b.arena.Reset()
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,126 @@
|
|||||||
|
package batch
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeBatchWriter struct {
|
||||||
|
bufs [][]byte
|
||||||
|
addrs []netip.AddrPort
|
||||||
|
ecns []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *fakeBatchWriter) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, ecns []byte) error {
|
||||||
|
// Snapshot — SendBatch.Flush nils its slot pointers right after WriteBatch
|
||||||
|
// returns, so tests must capture data before that happens.
|
||||||
|
w.bufs = make([][]byte, len(bufs))
|
||||||
|
for i, b := range bufs {
|
||||||
|
cp := make([]byte, len(b))
|
||||||
|
copy(cp, b)
|
||||||
|
w.bufs[i] = cp
|
||||||
|
}
|
||||||
|
w.addrs = append(w.addrs[:0], addrs...)
|
||||||
|
w.ecns = append(w.ecns[:0], ecns...)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchReserveCommitFlush(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
b := NewSendBatch(fw, 4, util.NewArena(32))
|
||||||
|
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:4242")
|
||||||
|
for i := 0; i < 4; i++ {
|
||||||
|
slot := b.Reserve(32)
|
||||||
|
if cap(slot) != 32 {
|
||||||
|
t.Fatalf("slot %d: cap=%d want 32", i, cap(slot))
|
||||||
|
}
|
||||||
|
pkt := append(slot[:0], byte(i), byte(i+1), byte(i+2))
|
||||||
|
b.Commit(pkt, ap, 0)
|
||||||
|
}
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
if len(fw.bufs) != 4 {
|
||||||
|
t.Fatalf("WriteBatch got %d bufs want 4", len(fw.bufs))
|
||||||
|
}
|
||||||
|
for i, buf := range fw.bufs {
|
||||||
|
if len(buf) != 3 || buf[0] != byte(i) {
|
||||||
|
t.Errorf("buf %d: %x", i, buf)
|
||||||
|
}
|
||||||
|
if fw.addrs[i] != ap {
|
||||||
|
t.Errorf("addr %d: got %v want %v", i, fw.addrs[i], ap)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flush again with nothing committed — should be a no-op.
|
||||||
|
fw.bufs = nil
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("empty Flush: %v", err)
|
||||||
|
}
|
||||||
|
if fw.bufs != nil {
|
||||||
|
t.Fatalf("empty Flush triggered WriteBatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reuse after Flush.
|
||||||
|
slot := b.Reserve(32)
|
||||||
|
if cap(slot) != 32 {
|
||||||
|
t.Fatalf("after Flush Reserve wrong cap: %d", cap(slot))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchSlotsDoNotOverlap(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
b := NewSendBatch(fw, 3, util.NewArena(8))
|
||||||
|
ap := netip.MustParseAddrPort("10.0.0.1:80")
|
||||||
|
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
s := b.Reserve(8)
|
||||||
|
pkt := append(s[:0], byte(0xA0+i), byte(0xB0+i))
|
||||||
|
b.Commit(pkt, ap, 0)
|
||||||
|
}
|
||||||
|
if err := b.Flush(); err != nil {
|
||||||
|
t.Fatalf("Flush: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, buf := range fw.bufs {
|
||||||
|
if buf[0] != byte(0xA0+i) || buf[1] != byte(0xB0+i) {
|
||||||
|
t.Errorf("slot %d corrupted: %x", i, buf)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSendBatchGrowPreservesCommitted(t *testing.T) {
|
||||||
|
fw := &fakeBatchWriter{}
|
||||||
|
// Tiny initial backing forces a grow on the second Reserve.
|
||||||
|
b := NewSendBatch(fw, 1, util.NewArena(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])
|
||||||
|
}
|
||||||
|
}
|
||||||
+8
-2
@@ -4,15 +4,21 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// defaultBatchBufSize is the per-Queue scratch size for Read on backends
|
||||||
|
// that don't do TSO segmentation. 65535 covers any single IP packet.
|
||||||
|
const defaultBatchBufSize = 65535
|
||||||
|
|
||||||
type Device interface {
|
type Device interface {
|
||||||
io.ReadWriteCloser
|
io.Closer
|
||||||
Activate() error
|
Activate() error
|
||||||
Networks() []netip.Prefix
|
Networks() []netip.Prefix
|
||||||
Name() string
|
Name() string
|
||||||
RoutesFor(netip.Addr) routing.Gateways
|
RoutesFor(netip.Addr) routing.Gateways
|
||||||
SupportsMultiqueue() bool
|
SupportsMultiqueue() bool
|
||||||
NewMultiQueueReader() (io.ReadWriteCloser, error)
|
NewMultiQueueReader() error
|
||||||
|
Readers() []tio.Queue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,358 @@
|
|||||||
|
//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,
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
//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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -4,10 +4,11 @@ package overlaytest
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
// NoopTun is an overlay.Device that silently discards every read and write.
|
// NoopTun is an overlay.Device that silently discards every read and write.
|
||||||
@@ -15,6 +16,10 @@ import (
|
|||||||
// exercise the datapath.
|
// exercise the datapath.
|
||||||
type NoopTun struct{}
|
type NoopTun struct{}
|
||||||
|
|
||||||
|
func (NoopTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func (NoopTun) RoutesFor(addr netip.Addr) routing.Gateways {
|
func (NoopTun) RoutesFor(addr netip.Addr) routing.Gateways {
|
||||||
return routing.Gateways{}
|
return routing.Gateways{}
|
||||||
}
|
}
|
||||||
@@ -31,7 +36,7 @@ func (NoopTun) Name() string {
|
|||||||
return "noop"
|
return "noop"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Read([]byte) (int, error) {
|
func (NoopTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -43,8 +48,12 @@ func (NoopTun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (NoopTun) NewMultiQueueReader() error {
|
||||||
return nil, errors.New("unsupported")
|
return errors.New("unsupported")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (NoopTun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{NoopTun{}}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (NoopTun) Close() error {
|
func (NoopTun) Close() error {
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
type pollQueueSet struct {
|
||||||
|
pq []*Poll
|
||||||
|
// pqi is exactly the same as pq, but stored as the interface type
|
||||||
|
pqi []Queue
|
||||||
|
shutdownFd int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewPollQueueSet() (QueueSet, error) {
|
||||||
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &pollQueueSet{
|
||||||
|
pq: []*Poll{},
|
||||||
|
pqi: []Queue{},
|
||||||
|
shutdownFd: shutdownFd,
|
||||||
|
}
|
||||||
|
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) Queues() []Queue {
|
||||||
|
return c.pqi
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) Add(fd int) error {
|
||||||
|
x, err := newPoll(fd, c.shutdownFd)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c.pq = append(c.pq, x)
|
||||||
|
c.pqi = append(c.pqi, x)
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) wakeForShutdown() error {
|
||||||
|
var buf [8]byte
|
||||||
|
binary.NativeEndian.PutUint64(buf[:], 1)
|
||||||
|
_, err := unix.Write(c.shutdownFd, buf[:])
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *pollQueueSet) Close() error {
|
||||||
|
if c.shutdownFd < 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
errs := []error{}
|
||||||
|
|
||||||
|
if err := c.wakeForShutdown(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, x := range c.pq {
|
||||||
|
if err := x.Close(); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// All Polls reference shutdownFd in their pollfd arrays, so close it
|
||||||
|
// only after every Poll.Close has returned.
|
||||||
|
if err := unix.Close(c.shutdownFd); err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
}
|
||||||
|
c.shutdownFd = -1
|
||||||
|
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
)
|
||||||
|
|
||||||
|
// QueueSet holds one or many Queue objects and helps close them in an orderly way.
|
||||||
|
type QueueSet interface {
|
||||||
|
io.Closer
|
||||||
|
Queues() []Queue
|
||||||
|
|
||||||
|
// Add takes a tun fd, adds it to the set, and prepares it for use as a Queue.
|
||||||
|
Add(fd int) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Capabilities advertises which kernel offload features a Queue successfully negotiated.
|
||||||
|
// Callers consult this to decide which coalescers to wire onto the write path.
|
||||||
|
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
|
||||||
|
// read goroutine plus a single writer (see Write below).
|
||||||
|
type Queue interface {
|
||||||
|
io.Closer
|
||||||
|
|
||||||
|
// Read will read at least 1 packet from the tun (up to len(p)).
|
||||||
|
// mem will be used to provide the backing for each of p[n].Bytes.
|
||||||
|
// Callers should size mem and p to avoid exhausting mem before p.
|
||||||
|
// Returns the number of packets actually read, or error.
|
||||||
|
Read(p []wire.TunPacket, mem []byte) (int, error)
|
||||||
|
|
||||||
|
// Write emits a single packet on the plaintext (outside→inside)
|
||||||
|
// delivery path.
|
||||||
|
Write(p []byte) (int, error)
|
||||||
|
|
||||||
|
// Capabilities returns the Queue's negotiated offload capabilities,
|
||||||
|
// or the zero value when q does not advertise any.
|
||||||
|
Capabilities() Capabilities
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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 (
|
||||||
|
GSOProtoNone GSOProto = iota
|
||||||
|
GSOProtoTCP
|
||||||
|
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
|
||||||
|
// fragments, in a single vectored write (writev with a leading
|
||||||
|
// virtio_net_hdr). This lets the coalescer avoid copying payload bytes
|
||||||
|
// between the caller's decrypt buffer and the TUN. Backends without GSO
|
||||||
|
// support do not implement this interface and coalescing is skipped.
|
||||||
|
//
|
||||||
|
// hdr contains the IPv4/IPv6 header prefix (mutable - callers will have
|
||||||
|
// filled in total length and IP csum). transportHdr is the TCP or UDP
|
||||||
|
// header (mutable - the L4 checksum field must hold the pseudo-header
|
||||||
|
// partial, single-fold not inverted, per virtio NEEDS_CSUM semantics).
|
||||||
|
// pays are non-overlapping payload fragments whose concatenation is the
|
||||||
|
// full superpacket payload; they are read-only from the writer's
|
||||||
|
// perspective and must remain valid until the call returns. Every segment
|
||||||
|
// 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
|
||||||
|
// QueueCapabilities) for the per-protocol negotiated capability; an
|
||||||
|
// implementation of GSOWriter is necessary but not sufficient since USO
|
||||||
|
// may not have been negotiated even when TSO was.
|
||||||
|
type GSOWriter interface {
|
||||||
|
WriteGSO(hdr []byte, transportHdr []byte, pays [][]byte, proto GSOProto) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// SupportsGSO reports whether w implements GSOWriter and the underlying
|
||||||
|
// queue advertises the negotiated capability for `want`. A writer that
|
||||||
|
// implements GSOWriter but not CapsProvider is treated as permissive
|
||||||
|
// (used by tests and fakes that don't negotiate).
|
||||||
|
func SupportsGSO(w Queue, want GSOProto) (GSOWriter, bool) {
|
||||||
|
gw, ok := w.(GSOWriter)
|
||||||
|
if !ok {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
caps := w.Capabilities()
|
||||||
|
switch want {
|
||||||
|
case GSOProtoTCP:
|
||||||
|
return gw, caps.TSO
|
||||||
|
case GSOProtoUDP:
|
||||||
|
return gw, caps.USO
|
||||||
|
default:
|
||||||
|
return gw, false
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Poll struct {
|
||||||
|
fd int
|
||||||
|
|
||||||
|
readPoll [2]unix.PollFd
|
||||||
|
writePoll [2]unix.PollFd
|
||||||
|
writeLock sync.Mutex
|
||||||
|
closed atomic.Bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newPoll(fd int, shutdownFd int) (*Poll, error) {
|
||||||
|
if err := unix.SetNonblock(fd, true); err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, fmt.Errorf("failed to set Poll device as nonblocking: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := &Poll{
|
||||||
|
fd: fd,
|
||||||
|
readPoll: [2]unix.PollFd{
|
||||||
|
{Fd: int32(fd), Events: unix.POLLIN},
|
||||||
|
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||||
|
},
|
||||||
|
writePoll: [2]unix.PollFd{
|
||||||
|
{Fd: int32(fd), Events: unix.POLLOUT},
|
||||||
|
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
||||||
|
},
|
||||||
|
writeLock: sync.Mutex{},
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// blockOnRead waits until the Poll fd is readable or shutdown has been signaled.
|
||||||
|
// Returns os.ErrClosed if Close was called.
|
||||||
|
func (t *Poll) blockOnRead() error {
|
||||||
|
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||||
|
var err error
|
||||||
|
for {
|
||||||
|
_, err = unix.Poll(t.readPoll[:], -1)
|
||||||
|
if err != unix.EINTR {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tunEvents := t.readPoll[0].Revents
|
||||||
|
shutdownEvents := t.readPoll[1].Revents
|
||||||
|
t.readPoll[0].Revents = 0
|
||||||
|
t.readPoll[1].Revents = 0
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
if tunEvents&problemFlags != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) blockOnWrite() error {
|
||||||
|
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
||||||
|
var err error
|
||||||
|
for {
|
||||||
|
_, err = unix.Poll(t.writePoll[:], -1)
|
||||||
|
if err != unix.EINTR {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.writeLock.Lock()
|
||||||
|
tunEvents := t.writePoll[0].Revents
|
||||||
|
shutdownEvents := t.writePoll[1].Revents
|
||||||
|
t.writePoll[0].Revents = 0
|
||||||
|
t.writePoll[1].Revents = 0
|
||||||
|
t.writeLock.Unlock()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
if tunEvents&problemFlags != 0 {
|
||||||
|
return os.ErrClosed
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) readOne(to []byte) (int, error) {
|
||||||
|
for {
|
||||||
|
n, errno := unix.Read(t.fd, to)
|
||||||
|
if errno == nil {
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
switch errno {
|
||||||
|
case unix.EAGAIN:
|
||||||
|
if err := t.blockOnRead(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
case unix.EINTR:
|
||||||
|
// retry
|
||||||
|
case unix.EBADF:
|
||||||
|
return 0, os.ErrClosed
|
||||||
|
default:
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Write(from []byte) (int, error) {
|
||||||
|
for {
|
||||||
|
n, errno := unix.Write(t.fd, from)
|
||||||
|
if errno == nil {
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
switch errno {
|
||||||
|
case unix.EAGAIN:
|
||||||
|
if err := t.blockOnWrite(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
case unix.EINTR:
|
||||||
|
// retry
|
||||||
|
case unix.EBADF:
|
||||||
|
return 0, os.ErrClosed
|
||||||
|
default:
|
||||||
|
return 0, errno
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Close() error {
|
||||||
|
if t.closed.Swap(true) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
//shutdownFd is owned by the container, so we should not close it
|
||||||
|
var err error
|
||||||
|
if t.fd >= 0 {
|
||||||
|
err = unix.Close(t.fd)
|
||||||
|
t.fd = -1
|
||||||
|
}
|
||||||
|
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *Poll) Capabilities() Capabilities {
|
||||||
|
return Capabilities{}
|
||||||
|
}
|
||||||
@@ -0,0 +1,106 @@
|
|||||||
|
//go:build linux && !android && !e2e_testing
|
||||||
|
// +build linux,!android,!e2e_testing
|
||||||
|
|
||||||
|
package tio
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/sys/unix"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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).
|
||||||
|
func newReadPipe(t *testing.T) int {
|
||||||
|
t.Helper()
|
||||||
|
var fds [2]int
|
||||||
|
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
||||||
|
t.Fatalf("pipe2: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
||||||
|
return fds[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPoll_WakeForShutdown_WakesFriends(t *testing.T) {
|
||||||
|
parent, err := NewPollQueueSet()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, parent.Add(newReadPipe(t)))
|
||||||
|
require.NoError(t, parent.Add(newReadPipe(t)))
|
||||||
|
// QueueSet.Close owns the read fds we Added — don't register a separate
|
||||||
|
// Cleanup to close them or we'll double-close whatever fd the kernel
|
||||||
|
// has since reused.
|
||||||
|
|
||||||
|
readers := parent.Queues()
|
||||||
|
errs := make([]error, len(readers))
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i, r := range readers {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(i int, r Queue) {
|
||||||
|
defer wg.Done()
|
||||||
|
pkts := make([]wire.TunPacket, 1)
|
||||||
|
_, errs[i] = r.Read(pkts, make([]byte, 64))
|
||||||
|
}(i, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
if err := parent.Close(); err != nil {
|
||||||
|
t.Fatalf("Close: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() { wg.Wait(); close(done) }()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("readers did not wake")
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, err := range errs {
|
||||||
|
if !errors.Is(err, os.ErrClosed) {
|
||||||
|
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPoll_Close_Idempotent(t *testing.T) {
|
||||||
|
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = unix.Close(shutdownFd) })
|
||||||
|
|
||||||
|
tf, err := newPoll(newReadPipe(t), shutdownFd)
|
||||||
|
require.NoError(t, err)
|
||||||
|
if err := tf.Close(); err != nil {
|
||||||
|
t.Fatalf("first Close: %v", err)
|
||||||
|
}
|
||||||
|
if err := tf.Close(); err != nil {
|
||||||
|
t.Fatalf("second Close should be a no-op, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPollQueueSet_Close_ClosesEventfd(t *testing.T) {
|
||||||
|
qs, err := NewPollQueueSet()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, qs.Add(newReadPipe(t)))
|
||||||
|
|
||||||
|
fd := qs.(*pollQueueSet).shutdownFd
|
||||||
|
require.NoError(t, qs.Close())
|
||||||
|
|
||||||
|
// Closing the eventfd again should fail with EBADF, proving Close
|
||||||
|
// actually released it.
|
||||||
|
if err := unix.Close(fd); err == nil {
|
||||||
|
t.Fatalf("eventfd %d still open after QueueSet.Close", fd)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Second Close must be a no-op (and must not double-close the eventfd
|
||||||
|
// in case the kernel handed it out to another caller in the meantime).
|
||||||
|
if err := qs.Close(); err != nil {
|
||||||
|
t.Fatalf("second Close: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+39
-8
@@ -13,12 +13,14 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
rwc io.ReadWriteCloser
|
||||||
fd int
|
fd int
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
@@ -26,16 +28,37 @@ type tun struct {
|
|||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.rwc.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Write(p []byte) (int, error) {
|
||||||
|
return t.rwc.Write(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Close() error {
|
||||||
|
return t.rwc.Close()
|
||||||
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
// XXX Android returns an fd in non-blocking mode which is necessary for shutdown to work properly.
|
// XXX Android returns an fd in non-blocking mode which is necessary for shutdown to work properly.
|
||||||
// Be sure not to call file.Fd() as it will set the fd to blocking mode.
|
// Be sure not to call file.Fd() as it will set the fd to blocking mode.
|
||||||
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/net/tun")
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
ReadWriteCloser: file,
|
rwc: file,
|
||||||
fd: deviceFd,
|
fd: deviceFd,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
@@ -62,7 +85,7 @@ func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
|||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t tun) Activate() error {
|
func (t *tun) Activate() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -99,6 +122,14 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for android")
|
return fmt.Errorf("TODO: multiqueue not implemented for android")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
//go:build (amd64 || arm64) && !e2e_testing
|
||||||
|
// +build amd64 arm64
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log/slog"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/wfp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// installInterfaceBypass installs a WFP PERMIT filter scoped to the wintun interface LUID so inbound traffic on the
|
||||||
|
// nebula adapter bypasses Windows Defender Firewall.
|
||||||
|
func installInterfaceBypass(l *slog.Logger, luid uint64) closer {
|
||||||
|
s, err := wfp.PermitInterface(luid)
|
||||||
|
if err != nil {
|
||||||
|
l.Warn("Failed to install WFP bypass filters on nebula interface", "error", err)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
l.Info("Installed WFP filters bypassing Windows Defender Firewall on nebula interface")
|
||||||
|
return s
|
||||||
|
}
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
//go:build !e2e_testing
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package overlay
|
||||||
|
|
||||||
|
import "log/slog"
|
||||||
|
|
||||||
|
// installInterfaceBypass is a no-op on windows-386 because we don't currently build for it.
|
||||||
|
func installInterfaceBypass(_ *slog.Logger, _ uint64) closer {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
+32
-18
@@ -16,14 +16,16 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
rwc io.ReadWriteCloser
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
DefaultMTU int
|
DefaultMTU int
|
||||||
@@ -124,11 +126,11 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*t
|
|||||||
}
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
ReadWriteCloser: os.NewFile(uintptr(fd), ""),
|
rwc: os.NewFile(uintptr(fd), ""),
|
||||||
Device: name,
|
Device: name,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
DefaultMTU: c.GetInt("tun.mtu", DefaultMTU),
|
DefaultMTU: c.GetInt("tun.mtu", DefaultMTU),
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err = t.reload(c, true)
|
err = t.reload(c, true)
|
||||||
@@ -158,8 +160,8 @@ func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, e
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
if t.ReadWriteCloser != nil {
|
if t.rwc != nil {
|
||||||
return t.ReadWriteCloser.Close()
|
return t.rwc.Close()
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -502,13 +504,17 @@ func delRoute(prefix netip.Prefix, gateway netroute.Addr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
buf := make([]byte, len(to)+4)
|
if len(p) == 0 || len(mem) <= 4 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
n, err := t.ReadWriteCloser.Read(buf)
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
copy(to, buf[4:])
|
n, err := t.rwc.Read(mem)
|
||||||
return n - 4, err
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[4:n]
|
||||||
|
return 1, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
// Write is only valid for single threaded use
|
||||||
@@ -536,7 +542,7 @@ func (t *tun) Write(from []byte) (int, error) {
|
|||||||
|
|
||||||
copy(buf[4:], from)
|
copy(buf[4:], from)
|
||||||
|
|
||||||
n, err := t.ReadWriteCloser.Write(buf)
|
n, err := t.rwc.Write(buf)
|
||||||
return n - 4, err
|
return n - 4, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -552,6 +558,14 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
return fmt.Errorf("TODO: multiqueue not implemented for darwin")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|||||||
+36
-6
@@ -10,7 +10,9 @@ import (
|
|||||||
|
|
||||||
"github.com/rcrowley/go-metrics"
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/iputil"
|
"github.com/slackhq/nebula/iputil"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type disabledTun struct {
|
type disabledTun struct {
|
||||||
@@ -18,9 +20,10 @@ type disabledTun struct {
|
|||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
|
|
||||||
// Track these metrics since we don't have the tun device to do it for us
|
// Track these metrics since we don't have the tun device to do it for us
|
||||||
tx metrics.Counter
|
tx metrics.Counter
|
||||||
rx metrics.Counter
|
rx metrics.Counter
|
||||||
l *slog.Logger
|
numReaders int
|
||||||
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled bool, l *slog.Logger) *disabledTun {
|
||||||
@@ -28,6 +31,7 @@ func newDisabledTun(vpnNetworks []netip.Prefix, queueLen int, metricsEnabled boo
|
|||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
read: make(chan []byte, queueLen),
|
read: make(chan []byte, queueLen),
|
||||||
l: l,
|
l: l,
|
||||||
|
numReaders: 1,
|
||||||
}
|
}
|
||||||
|
|
||||||
if metricsEnabled {
|
if metricsEnabled {
|
||||||
@@ -57,7 +61,7 @@ func (*disabledTun) Name() string {
|
|||||||
return "disabled"
|
return "disabled"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) Read(b []byte) (int, error) {
|
func (t *disabledTun) readOne(b []byte) (int, error) {
|
||||||
r, ok := <-t.read
|
r, ok := <-t.read
|
||||||
if !ok {
|
if !ok {
|
||||||
return 0, io.EOF
|
return 0, io.EOF
|
||||||
@@ -75,6 +79,19 @@ func (t *disabledTun) Read(b []byte) (int, error) {
|
|||||||
return copy(b, r), nil
|
return copy(b, r), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
func (t *disabledTun) handleICMPEchoRequest(b []byte) bool {
|
||||||
out := make([]byte, len(b))
|
out := make([]byte, len(b))
|
||||||
out = iputil.CreateICMPEchoResponse(b, out)
|
out = iputil.CreateICMPEchoResponse(b, out)
|
||||||
@@ -110,8 +127,21 @@ func (t *disabledTun) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *disabledTun) NewMultiQueueReader() error {
|
||||||
return t, nil
|
t.numReaders++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Readers() []tio.Queue {
|
||||||
|
out := make([]tio.Queue, t.numReaders)
|
||||||
|
for i := range t.numReaders {
|
||||||
|
out[i] = t
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *disabledTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *disabledTun) Close() error {
|
func (t *disabledTun) Close() error {
|
||||||
|
|||||||
@@ -1,120 +0,0 @@
|
|||||||
//go:build linux && !android && !e2e_testing
|
|
||||||
// +build linux,!android,!e2e_testing
|
|
||||||
|
|
||||||
package overlay
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"os"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
|
||||||
)
|
|
||||||
|
|
||||||
// newReadPipe returns a read fd. The matching write fd is registered for cleanup.
|
|
||||||
// The caller takes ownership of the read fd (pass it to newTunFd / newFriend).
|
|
||||||
func newReadPipe(t *testing.T) int {
|
|
||||||
t.Helper()
|
|
||||||
var fds [2]int
|
|
||||||
if err := unix.Pipe2(fds[:], unix.O_CLOEXEC); err != nil {
|
|
||||||
t.Fatalf("pipe2: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() { _ = unix.Close(fds[1]) })
|
|
||||||
return fds[0]
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTunFile_WakeForShutdown_UnblocksRead(t *testing.T) {
|
|
||||||
tf, err := newTunFd(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTunFd: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() { _ = tf.Close() })
|
|
||||||
|
|
||||||
done := make(chan error, 1)
|
|
||||||
go func() {
|
|
||||||
_, err := tf.Read(make([]byte, 64))
|
|
||||||
done <- err
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Verify Read is actually blocked in poll.
|
|
||||||
select {
|
|
||||||
case err := <-done:
|
|
||||||
t.Fatalf("Read returned before shutdown signal: %v", err)
|
|
||||||
case <-time.After(50 * time.Millisecond):
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := tf.wakeForShutdown(); err != nil {
|
|
||||||
t.Fatalf("wakeForShutdown: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case err := <-done:
|
|
||||||
if !errors.Is(err, os.ErrClosed) {
|
|
||||||
t.Fatalf("expected os.ErrClosed, got %v", err)
|
|
||||||
}
|
|
||||||
case <-time.After(2 * time.Second):
|
|
||||||
t.Fatal("Read did not wake on shutdown")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTunFile_WakeForShutdown_WakesFriends(t *testing.T) {
|
|
||||||
parent, err := newTunFd(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTunFd: %v", err)
|
|
||||||
}
|
|
||||||
friend, err := parent.newFriend(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
_ = parent.Close()
|
|
||||||
t.Fatalf("newFriend: %v", err)
|
|
||||||
}
|
|
||||||
t.Cleanup(func() {
|
|
||||||
_ = friend.Close()
|
|
||||||
_ = parent.Close()
|
|
||||||
})
|
|
||||||
|
|
||||||
readers := []*tunFile{parent, friend}
|
|
||||||
errs := make([]error, len(readers))
|
|
||||||
var wg sync.WaitGroup
|
|
||||||
for i, r := range readers {
|
|
||||||
wg.Add(1)
|
|
||||||
go func(i int, r *tunFile) {
|
|
||||||
defer wg.Done()
|
|
||||||
_, errs[i] = r.Read(make([]byte, 64))
|
|
||||||
}(i, r)
|
|
||||||
}
|
|
||||||
|
|
||||||
time.Sleep(50 * time.Millisecond)
|
|
||||||
|
|
||||||
if err := parent.wakeForShutdown(); err != nil {
|
|
||||||
t.Fatalf("wakeForShutdown: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
done := make(chan struct{})
|
|
||||||
go func() { wg.Wait(); close(done) }()
|
|
||||||
select {
|
|
||||||
case <-done:
|
|
||||||
case <-time.After(2 * time.Second):
|
|
||||||
t.Fatal("readers did not wake")
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, err := range errs {
|
|
||||||
if !errors.Is(err, os.ErrClosed) {
|
|
||||||
t.Errorf("reader %d: expected os.ErrClosed, got %v", i, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTunFile_Close_Idempotent(t *testing.T) {
|
|
||||||
tf, err := newTunFd(newReadPipe(t))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("newTunFd: %v", err)
|
|
||||||
}
|
|
||||||
if err := tf.Close(); err != nil {
|
|
||||||
t.Fatalf("first Close: %v", err)
|
|
||||||
}
|
|
||||||
if err := tf.Close(); err != nil {
|
|
||||||
t.Fatalf("second Close should be a no-op, got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+26
-5
@@ -7,7 +7,6 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"io/fs"
|
"io/fs"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -18,9 +17,10 @@ import (
|
|||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
@@ -157,7 +157,20 @@ func (t *tun) blockOnWrite() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) readOne(to []byte) (int, error) {
|
||||||
// first 4 bytes is protocol family, in network byte order
|
// first 4 bytes is protocol family, in network byte order
|
||||||
var head [4]byte
|
var head [4]byte
|
||||||
iovecs := [2]syscall.Iovec{
|
iovecs := [2]syscall.Iovec{
|
||||||
@@ -565,8 +578,8 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
return fmt.Errorf("TODO: multiqueue not implemented for freebsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
@@ -593,6 +606,14 @@ func (t *tun) addRoutes(logErrors bool) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func (t *tun) removeRoutes(routes []Route) error {
|
func (t *tun) removeRoutes(routes []Route) error {
|
||||||
for _, r := range routes {
|
for _, r := range routes {
|
||||||
if !r.Install {
|
if !r.Install {
|
||||||
|
|||||||
+39
-17
@@ -16,18 +16,41 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
io.ReadWriteCloser
|
rwc io.ReadWriteCloser
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) <= 4 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.rwc.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[4:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Write(p []byte) (int, error) {
|
||||||
|
return t.rwc.Write(p)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Close() error {
|
||||||
|
return t.rwc.Close()
|
||||||
|
}
|
||||||
|
|
||||||
func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error) {
|
func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error) {
|
||||||
return nil, fmt.Errorf("newTun not supported in iOS")
|
return nil, fmt.Errorf("newTun not supported in iOS")
|
||||||
}
|
}
|
||||||
@@ -35,9 +58,9 @@ func newTun(_ *config.C, _ *slog.Logger, _ []netip.Prefix, _ bool) (*tun, error)
|
|||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
file := os.NewFile(uintptr(deviceFd), "/dev/tun")
|
||||||
t := &tun{
|
t := &tun{
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
ReadWriteCloser: &tunReadCloser{f: file},
|
rwc: &tunReadCloser{f: file},
|
||||||
l: l,
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err := t.reload(c, true)
|
err := t.reload(c, true)
|
||||||
@@ -96,18 +119,9 @@ type tunReadCloser struct {
|
|||||||
wBuf []byte
|
wBuf []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Read returns a packet with the BSD 4-byte header, watch out!
|
||||||
func (tr *tunReadCloser) Read(to []byte) (int, error) {
|
func (tr *tunReadCloser) Read(to []byte) (int, error) {
|
||||||
tr.rMu.Lock()
|
return tr.f.Read(to)
|
||||||
defer tr.rMu.Unlock()
|
|
||||||
|
|
||||||
if cap(tr.rBuf) < len(to)+4 {
|
|
||||||
tr.rBuf = make([]byte, len(to)+4)
|
|
||||||
}
|
|
||||||
tr.rBuf = tr.rBuf[:len(to)+4]
|
|
||||||
|
|
||||||
n, err := tr.f.Read(tr.rBuf)
|
|
||||||
copy(to, tr.rBuf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (tr *tunReadCloser) Write(from []byte) (int, error) {
|
func (tr *tunReadCloser) Write(from []byte) (int, error) {
|
||||||
@@ -155,6 +169,14 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for ios")
|
return fmt.Errorf("TODO: multiqueue not implemented for ios")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|||||||
+84
-241
@@ -4,9 +4,7 @@
|
|||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/binary"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
@@ -19,180 +17,15 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"github.com/vishvananda/netlink"
|
"github.com/vishvananda/netlink"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
// tunFile 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.
|
|
||||||
type tunFile struct {
|
|
||||||
fd int
|
|
||||||
shutdownFd int
|
|
||||||
lastOne bool
|
|
||||||
readPoll [2]unix.PollFd
|
|
||||||
writePoll [2]unix.PollFd
|
|
||||||
closed bool
|
|
||||||
}
|
|
||||||
|
|
||||||
// newFriend makes a tunFile for a MultiQueueReader that copies the shutdown eventfd from the parent tun
|
|
||||||
func (r *tunFile) newFriend(fd int) (*tunFile, error) {
|
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
|
||||||
}
|
|
||||||
return &tunFile{
|
|
||||||
fd: fd,
|
|
||||||
shutdownFd: r.shutdownFd,
|
|
||||||
readPoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLIN},
|
|
||||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
writePoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
|
||||||
{Fd: int32(r.shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func newTunFd(fd int) (*tunFile, error) {
|
|
||||||
if err := unix.SetNonblock(fd, true); err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to set tun fd non-blocking: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
shutdownFd, err := unix.Eventfd(0, unix.EFD_NONBLOCK|unix.EFD_CLOEXEC)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create eventfd: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
out := &tunFile{
|
|
||||||
fd: fd,
|
|
||||||
shutdownFd: shutdownFd,
|
|
||||||
lastOne: true,
|
|
||||||
readPoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLIN},
|
|
||||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
writePoll: [2]unix.PollFd{
|
|
||||||
{Fd: int32(fd), Events: unix.POLLOUT},
|
|
||||||
{Fd: int32(shutdownFd), Events: unix.POLLIN},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) blockOnRead() error {
|
|
||||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
|
||||||
var err error
|
|
||||||
for {
|
|
||||||
_, err = unix.Poll(r.readPoll[:], -1)
|
|
||||||
if err != unix.EINTR {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
//always reset these!
|
|
||||||
tunEvents := r.readPoll[0].Revents
|
|
||||||
shutdownEvents := r.readPoll[1].Revents
|
|
||||||
r.readPoll[0].Revents = 0
|
|
||||||
r.readPoll[1].Revents = 0
|
|
||||||
//do the err check before trusting the potentially bogus bits we just got
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
} else if tunEvents&problemFlags != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) blockOnWrite() error {
|
|
||||||
const problemFlags = unix.POLLHUP | unix.POLLNVAL | unix.POLLERR
|
|
||||||
var err error
|
|
||||||
for {
|
|
||||||
_, err = unix.Poll(r.writePoll[:], -1)
|
|
||||||
if err != unix.EINTR {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
//always reset these!
|
|
||||||
tunEvents := r.writePoll[0].Revents
|
|
||||||
shutdownEvents := r.writePoll[1].Revents
|
|
||||||
r.writePoll[0].Revents = 0
|
|
||||||
r.writePoll[1].Revents = 0
|
|
||||||
//do the err check before trusting the potentially bogus bits we just got
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
if shutdownEvents&(unix.POLLIN|problemFlags) != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
} else if tunEvents&problemFlags != 0 {
|
|
||||||
return os.ErrClosed
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) Read(buf []byte) (int, error) {
|
|
||||||
for {
|
|
||||||
if n, err := unix.Read(r.fd, buf); err == nil {
|
|
||||||
return n, nil
|
|
||||||
} else if err == unix.EAGAIN {
|
|
||||||
if err = r.blockOnRead(); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
} else if err == unix.EINTR {
|
|
||||||
continue
|
|
||||||
} else if err == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
} else {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) Write(buf []byte) (int, error) {
|
|
||||||
for {
|
|
||||||
if n, err := unix.Write(r.fd, buf); err == nil {
|
|
||||||
return n, nil
|
|
||||||
} else if err == unix.EAGAIN {
|
|
||||||
if err = r.blockOnWrite(); err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
} else if err == unix.EINTR {
|
|
||||||
continue
|
|
||||||
} else if err == unix.EBADF {
|
|
||||||
return 0, os.ErrClosed
|
|
||||||
} else {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) wakeForShutdown() error {
|
|
||||||
var buf [8]byte
|
|
||||||
binary.NativeEndian.PutUint64(buf[:], 1)
|
|
||||||
_, err := unix.Write(int(r.readPoll[1].Fd), buf[:])
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r *tunFile) Close() error {
|
|
||||||
if r.closed { // avoid closing more than once. Technically a fd could get re-used, which would be a problem
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
r.closed = true
|
|
||||||
if r.lastOne {
|
|
||||||
_ = unix.Close(r.shutdownFd)
|
|
||||||
}
|
|
||||||
return unix.Close(r.fd)
|
|
||||||
}
|
|
||||||
|
|
||||||
type tun struct {
|
type tun struct {
|
||||||
*tunFile
|
readers tio.QueueSet
|
||||||
readers []*tunFile
|
|
||||||
closeLock sync.Mutex
|
closeLock sync.Mutex
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
@@ -239,7 +72,9 @@ type ifreqQLEN struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
t, err := newTunGeneric(c, l, deviceFd, vpnNetworks)
|
// We don't know what flags the caller opened this fd with and can't turn
|
||||||
|
// on IFF_VNET_HDR after TUNSETIFF, so skip offload on inherited fds.
|
||||||
|
t, err := newTunGeneric(c, l, deviceFd, false, false, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -249,46 +84,60 @@ func newTunFromFd(c *config.C, l *slog.Logger, deviceFd int, vpnNetworks []netip
|
|||||||
return t, nil
|
return t, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
// openTunDev opens /dev/net/tun, creating the device node first if it's
|
||||||
|
// missing (docker containers occasionally omit it).
|
||||||
|
func openTunDev() (int, error) {
|
||||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if err != nil {
|
if err == nil {
|
||||||
// If /dev/net/tun doesn't exist, try to create it (will happen in docker)
|
return fd, nil
|
||||||
if os.IsNotExist(err) {
|
|
||||||
err = os.MkdirAll("/dev/net", 0755)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
|
||||||
}
|
|
||||||
err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200)))
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
if !os.IsNotExist(err) {
|
||||||
|
return -1, err
|
||||||
|
}
|
||||||
|
if err = os.MkdirAll("/dev/net", 0755); err != nil {
|
||||||
|
return -1, fmt.Errorf("/dev/net/tun doesn't exist, failed to mkdir -p /dev/net: %w", err)
|
||||||
|
}
|
||||||
|
if err = unix.Mknod("/dev/net/tun", unix.S_IFCHR|0600, int(unix.Mkdev(10, 200))); err != nil {
|
||||||
|
return -1, fmt.Errorf("failed to create /dev/net/tun: %w", err)
|
||||||
|
}
|
||||||
|
fd, err = unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
|
if err != nil {
|
||||||
|
return -1, fmt.Errorf("created /dev/net/tun, but still failed: %w", err)
|
||||||
|
}
|
||||||
|
return fd, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// tunSetIff runs TUNSETIFF with the given flags and returns the kernel-chosen
|
||||||
|
// device name on success.
|
||||||
|
func tunSetIff(fd int, name string, flags uint16) (string, error) {
|
||||||
var req ifReq
|
var req ifReq
|
||||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
req.Flags = flags
|
||||||
|
copy(req.Name[:], name)
|
||||||
|
if err := ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return strings.Trim(string(req.Name[:]), "\x00"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue bool) (*tun, error) {
|
||||||
|
baseFlags := uint16(unix.IFF_TUN | unix.IFF_NO_PI)
|
||||||
if multiqueue {
|
if multiqueue {
|
||||||
req.Flags |= unix.IFF_MULTI_QUEUE
|
baseFlags |= unix.IFF_MULTI_QUEUE
|
||||||
}
|
}
|
||||||
nameStr := c.GetString("tun.dev", "")
|
nameStr := c.GetString("tun.dev", "")
|
||||||
copy(req.Name[:], nameStr)
|
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
|
||||||
return nil, &NameError{
|
|
||||||
Name: nameStr,
|
|
||||||
Underlying: err,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
name := strings.Trim(string(req.Name[:]), "\x00")
|
|
||||||
|
|
||||||
t, err := newTunGeneric(c, l, fd, vpnNetworks)
|
fd, err := openTunDev()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
name, err := tunSetIff(fd, nameStr, baseFlags)
|
||||||
|
if err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, &NameError{Name: nameStr, Underlying: err}
|
||||||
|
}
|
||||||
|
|
||||||
|
t, err := newTunGeneric(c, l, fd, false, false, vpnNetworks)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -299,15 +148,21 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, multiqueue
|
|||||||
}
|
}
|
||||||
|
|
||||||
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
// newTunGeneric does all the stuff common to different tun initialization paths. It will close your files on error.
|
||||||
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vpnNetworks []netip.Prefix) (*tun, error) {
|
func newTunGeneric(c *config.C, l *slog.Logger, fd int, vnetHdr, usoEnabled bool, vpnNetworks []netip.Prefix) (*tun, error) {
|
||||||
tfd, err := newTunFd(fd)
|
qs, err := tio.NewPollQueueSet()
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
err = qs.Add(fd)
|
||||||
|
if err != nil {
|
||||||
|
_ = unix.Close(fd)
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
t := &tun{
|
t := &tun{
|
||||||
tunFile: tfd,
|
readers: qs,
|
||||||
readers: []*tunFile{tfd},
|
|
||||||
closeLock: sync.Mutex{},
|
closeLock: sync.Mutex{},
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
TXQueueLen: c.GetInt("tun.tx_queue", 500),
|
||||||
@@ -410,32 +265,29 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
t.closeLock.Lock()
|
t.closeLock.Lock()
|
||||||
defer t.closeLock.Unlock()
|
defer t.closeLock.Unlock()
|
||||||
|
|
||||||
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
fd, err := unix.Open("/dev/net/tun", os.O_RDWR, 0)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var req ifReq
|
flags := uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
||||||
req.Flags = uint16(unix.IFF_TUN | unix.IFF_NO_PI | unix.IFF_MULTI_QUEUE)
|
|
||||||
copy(req.Name[:], t.Device)
|
if _, err = tunSetIff(fd, t.Device, flags); err != nil {
|
||||||
if err = ioctl(uintptr(fd), uintptr(unix.TUNSETIFF), uintptr(unsafe.Pointer(&req))); err != nil {
|
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
out, err := t.tunFile.newFriend(fd)
|
err = t.readers.Add(fd)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
_ = unix.Close(fd)
|
_ = unix.Close(fd)
|
||||||
return nil, err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
t.readers = append(t.readers, out)
|
return nil
|
||||||
|
|
||||||
return out, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
func (t *tun) RoutesFor(ip netip.Addr) routing.Gateways {
|
||||||
@@ -603,6 +455,15 @@ func (t *tun) setDefaultRoute(cidr netip.Prefix) error {
|
|||||||
Table: unix.RT_TABLE_MAIN,
|
Table: unix.RT_TABLE_MAIN,
|
||||||
Type: unix.RTN_UNICAST,
|
Type: unix.RTN_UNICAST,
|
||||||
}
|
}
|
||||||
|
// Match the metric the kernel uses for its auto-installed connected
|
||||||
|
// route, so RouteReplace overwrites it in place instead of adding a
|
||||||
|
// second route at a worse metric. IPv6 connected routes are installed
|
||||||
|
// at metric 256 (IP6_RT_PRIO_KERN); IPv4 uses 0. Without this, the
|
||||||
|
// kernel route wins lookups and our MTU / AdvMSS / Features never
|
||||||
|
// apply on v6.
|
||||||
|
if cidr.Addr().Is6() {
|
||||||
|
nr.Priority = 256
|
||||||
|
}
|
||||||
err := netlink.RouteReplace(&nr)
|
err := netlink.RouteReplace(&nr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
t.l.Warn("Failed to set default route MTU, retrying", "error", err, "cidr", cidr)
|
||||||
@@ -869,6 +730,10 @@ func (t *tun) updateRoutes(r netlink.RouteUpdate) {
|
|||||||
t.routeTree.Store(newTree)
|
t.routeTree.Store(newTree)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return t.readers.Queues()
|
||||||
|
}
|
||||||
|
|
||||||
func (t *tun) Close() error {
|
func (t *tun) Close() error {
|
||||||
t.closeLock.Lock()
|
t.closeLock.Lock()
|
||||||
defer t.closeLock.Unlock()
|
defer t.closeLock.Unlock()
|
||||||
@@ -878,32 +743,10 @@ func (t *tun) Close() error {
|
|||||||
t.routeChan = nil
|
t.routeChan = nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Signal all readers blocked in poll to wake up and exit
|
|
||||||
_ = t.tunFile.wakeForShutdown()
|
|
||||||
|
|
||||||
if t.ioctlFd > 0 {
|
if t.ioctlFd > 0 {
|
||||||
_ = unix.Close(int(t.ioctlFd))
|
_ = unix.Close(int(t.ioctlFd))
|
||||||
t.ioctlFd = 0
|
t.ioctlFd = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
for i := range t.readers {
|
return t.readers.Close()
|
||||||
if i == 0 {
|
|
||||||
continue //we want to close the zeroth reader last
|
|
||||||
}
|
|
||||||
err := t.readers[i].Close()
|
|
||||||
if err != nil {
|
|
||||||
t.l.Error("error closing tun reader", "reader", i, "error", err)
|
|
||||||
} else {
|
|
||||||
t.l.Info("closed tun reader", "reader", i)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
//this is t.readers[0] too
|
|
||||||
err := t.tunFile.Close()
|
|
||||||
if err != nil {
|
|
||||||
t.l.Error("error closing tun reader", "reader", 0, "error", err)
|
|
||||||
} else {
|
|
||||||
t.l.Info("closed tun reader", "reader", 0)
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,7 +3,9 @@
|
|||||||
|
|
||||||
package overlay
|
package overlay
|
||||||
|
|
||||||
import "testing"
|
import (
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
var runAdvMSSTests = []struct {
|
var runAdvMSSTests = []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
+26
-4
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,8 +16,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
@@ -68,6 +69,27 @@ type tun struct {
|
|||||||
fd int
|
fd int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.readOne(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
||||||
@@ -141,7 +163,7 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
func (t *tun) readOne(to []byte) (int, error) {
|
||||||
rc, err := t.f.SyscallConn()
|
rc, err := t.f.SyscallConn()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("failed to get syscall conn for tun: %w", err)
|
return 0, fmt.Errorf("failed to get syscall conn for tun: %w", err)
|
||||||
@@ -394,8 +416,8 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
return fmt.Errorf("TODO: multiqueue not implemented for netbsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
|
|||||||
+25
-12
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -17,8 +16,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
netroute "golang.org/x/net/route"
|
netroute "golang.org/x/net/route"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
@@ -61,6 +62,19 @@ type tun struct {
|
|||||||
out []byte
|
out []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) <= 4 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.f.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[4:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
var deviceNameRE = regexp.MustCompile(`^tun[0-9]+$`)
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (*tun, error) {
|
||||||
@@ -124,15 +138,6 @@ func (t *tun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) Read(to []byte) (int, error) {
|
|
||||||
buf := make([]byte, len(to)+4)
|
|
||||||
|
|
||||||
n, err := t.f.Read(buf)
|
|
||||||
|
|
||||||
copy(to, buf[4:])
|
|
||||||
return n - 4, err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write is only valid for single threaded use
|
// Write is only valid for single threaded use
|
||||||
func (t *tun) Write(from []byte) (int, error) {
|
func (t *tun) Write(from []byte) (int, error) {
|
||||||
buf := t.out
|
buf := t.out
|
||||||
@@ -314,8 +319,8 @@ func (t *tun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *tun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
return fmt.Errorf("TODO: multiqueue not implemented for openbsd")
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *tun) addRoutes(logErrors bool) error {
|
func (t *tun) addRoutes(logErrors bool) error {
|
||||||
@@ -366,6 +371,14 @@ func (t *tun) deviceBytes() (o [16]byte) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *tun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *tun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func addRoute(prefix netip.Prefix, gateways []netip.Prefix) error {
|
func addRoute(prefix netip.Prefix, gateways []netip.Prefix) error {
|
||||||
sock, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
sock, err := unix.Socket(unix.AF_ROUTE, unix.SOCK_RAW, unix.AF_UNSPEC)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+26
-3
@@ -14,8 +14,10 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/udp"
|
"github.com/slackhq/nebula/udp"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
type TestTun struct {
|
type TestTun struct {
|
||||||
@@ -162,7 +164,20 @@ func (t *TestTun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) Read(b []byte) (int, error) {
|
func (t *TestTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) read(b []byte) (int, error) {
|
||||||
p, ok := <-t.rxPackets
|
p, ok := <-t.rxPackets
|
||||||
if !ok {
|
if !ok {
|
||||||
return 0, os.ErrClosed
|
return 0, os.ErrClosed
|
||||||
@@ -177,10 +192,18 @@ func (t *TestTun) Read(b []byte) (int, error) {
|
|||||||
return n, nil
|
return n, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *TestTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
func (t *TestTun) SupportsMultiqueue() bool {
|
func (t *TestTun) SupportsMultiqueue() bool {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *TestTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *TestTun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented")
|
return fmt.Errorf("TODO: multiqueue not implemented")
|
||||||
}
|
}
|
||||||
|
|||||||
+79
-23
@@ -6,7 +6,6 @@ package overlay
|
|||||||
import (
|
import (
|
||||||
"crypto"
|
"crypto"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"os"
|
"os"
|
||||||
@@ -18,26 +17,50 @@ import (
|
|||||||
|
|
||||||
"github.com/gaissmai/bart"
|
"github.com/gaissmai/bart"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
"github.com/slackhq/nebula/util"
|
"github.com/slackhq/nebula/util"
|
||||||
"github.com/slackhq/nebula/wintun"
|
"github.com/slackhq/nebula/wintun"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
"golang.org/x/sys/windows"
|
"golang.org/x/sys/windows"
|
||||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type closer interface {
|
||||||
|
Close()
|
||||||
|
}
|
||||||
|
|
||||||
const tunGUIDLabel = "Fixed Nebula Windows GUID v1"
|
const tunGUIDLabel = "Fixed Nebula Windows GUID v1"
|
||||||
|
|
||||||
type winTun struct {
|
type winTun struct {
|
||||||
Device string
|
Device string
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
MTU int
|
MTU int
|
||||||
Routes atomic.Pointer[[]Route]
|
Routes atomic.Pointer[[]Route]
|
||||||
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
routeTree atomic.Pointer[bart.Table[routing.Gateways]]
|
||||||
l *slog.Logger
|
guid windows.GUID
|
||||||
|
networkCategory networkCategory
|
||||||
|
setCategory bool
|
||||||
|
bypassWDF bool
|
||||||
|
wdfBypass closer
|
||||||
|
l *slog.Logger
|
||||||
|
|
||||||
tun *wintun.NativeTun
|
tun *wintun.NativeTun
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *winTun) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := t.tun.Read(mem, 0)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
func newTunFromFd(_ *config.C, _ *slog.Logger, _ int, _ []netip.Prefix) (Device, error) {
|
||||||
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
return nil, fmt.Errorf("newTunFromFd not supported in Windows")
|
||||||
}
|
}
|
||||||
@@ -54,11 +77,20 @@ func newTun(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, _ bool) (*w
|
|||||||
return nil, fmt.Errorf("generate GUID failed: %w", err)
|
return nil, fmt.Errorf("generate GUID failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
cat, setCat, err := parseNetworkCategory(c.GetString("tun.network_category", "private"))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
t := &winTun{
|
t := &winTun{
|
||||||
Device: deviceName,
|
Device: deviceName,
|
||||||
vpnNetworks: vpnNetworks,
|
vpnNetworks: vpnNetworks,
|
||||||
MTU: c.GetInt("tun.mtu", DefaultMTU),
|
MTU: c.GetInt("tun.mtu", DefaultMTU),
|
||||||
l: l,
|
guid: *guid,
|
||||||
|
networkCategory: cat,
|
||||||
|
setCategory: setCat,
|
||||||
|
bypassWDF: c.GetBool("tun.windows_bypass_wdf", true),
|
||||||
|
l: l,
|
||||||
}
|
}
|
||||||
|
|
||||||
err = t.reload(c, true)
|
err = t.reload(c, true)
|
||||||
@@ -142,6 +174,17 @@ func (t *winTun) Activate() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if t.setCategory {
|
||||||
|
// The wintun adapter takes a moment to register with the Network List
|
||||||
|
// Manager, so we apply the category in the background and retry until
|
||||||
|
// it shows up.
|
||||||
|
go applyNetworkCategory(t.l, t.guid, t.networkCategory)
|
||||||
|
}
|
||||||
|
|
||||||
|
if t.bypassWDF {
|
||||||
|
t.wdfBypass = installInterfaceBypass(t.l, uint64(t.tun.LUID()))
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -156,11 +199,8 @@ func (t *winTun) addRoutes(logErrors bool) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add our unsafe route
|
// Add our unsafe route as an on-link route to the nebula tun device.
|
||||||
// Windows does not support multipath routes natively, so we install only a single route.
|
err := luid.AddRoute(r.Cidr, unspecifiedNextHop(r.Cidr), uint32(r.Metric))
|
||||||
// This is not a problem as traffic will always be sent to Nebula which handles the multipath routing internally.
|
|
||||||
// In effect this provides multipath routing support to windows supporting loadbalancing and redundancy.
|
|
||||||
err := luid.AddRoute(r.Cidr, r.Via[0].Addr(), uint32(r.Metric))
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
retErr := util.NewContextualError("Failed to add route", map[string]any{"route": r}, err)
|
retErr := util.NewContextualError("Failed to add route", map[string]any{"route": r}, err)
|
||||||
if logErrors {
|
if logErrors {
|
||||||
@@ -206,7 +246,7 @@ func (t *winTun) removeRoutes(routes []Route) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// See comment on luid.AddRoute
|
// See comment on luid.AddRoute
|
||||||
err := luid.DeleteRoute(r.Cidr, r.Via[0].Addr())
|
err := luid.DeleteRoute(r.Cidr, unspecifiedNextHop(r.Cidr))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.l.Error("Failed to remove route", "error", err, "route", r)
|
t.l.Error("Failed to remove route", "error", err, "route", r)
|
||||||
} else {
|
} else {
|
||||||
@@ -229,10 +269,6 @@ func (t *winTun) Name() string {
|
|||||||
return t.Device
|
return t.Device
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) Read(b []byte) (int, error) {
|
|
||||||
return t.tun.Read(b, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (t *winTun) Write(b []byte) (int, error) {
|
func (t *winTun) Write(b []byte) (int, error) {
|
||||||
return t.tun.Write(b, 0)
|
return t.tun.Write(b, 0)
|
||||||
}
|
}
|
||||||
@@ -241,8 +277,16 @@ func (t *winTun) SupportsMultiqueue() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (t *winTun) NewMultiQueueReader() error {
|
||||||
return nil, fmt.Errorf("TODO: multiqueue not implemented for windows")
|
return fmt.Errorf("TODO: multiqueue not implemented for windows")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *winTun) Readers() []tio.Queue {
|
||||||
|
return []tio.Queue{t}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *winTun) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *winTun) Close() error {
|
func (t *winTun) Close() error {
|
||||||
@@ -258,9 +302,21 @@ func (t *winTun) Close() error {
|
|||||||
_ = luid.FlushDNS(windows.AF_INET)
|
_ = luid.FlushDNS(windows.AF_INET)
|
||||||
_ = luid.FlushDNS(windows.AF_INET6)
|
_ = luid.FlushDNS(windows.AF_INET6)
|
||||||
|
|
||||||
|
if t.wdfBypass != nil {
|
||||||
|
t.wdfBypass.Close()
|
||||||
|
t.wdfBypass = nil
|
||||||
|
}
|
||||||
|
|
||||||
return t.tun.Close()
|
return t.tun.Close()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func unspecifiedNextHop(p netip.Prefix) netip.Addr {
|
||||||
|
if p.Addr().Is4() {
|
||||||
|
return netip.IPv4Unspecified()
|
||||||
|
}
|
||||||
|
return netip.IPv6Unspecified()
|
||||||
|
}
|
||||||
|
|
||||||
func generateGUIDByDeviceName(name string) (*windows.GUID, error) {
|
func generateGUIDByDeviceName(name string) (*windows.GUID, error) {
|
||||||
// GUID is 128 bit
|
// GUID is 128 bit
|
||||||
hash := crypto.MD5.New()
|
hash := crypto.MD5.New()
|
||||||
|
|||||||
+33
-5
@@ -6,7 +6,9 @@ import (
|
|||||||
"net/netip"
|
"net/netip"
|
||||||
|
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/overlay/tio"
|
||||||
"github.com/slackhq/nebula/routing"
|
"github.com/slackhq/nebula/routing"
|
||||||
|
"github.com/slackhq/nebula/wire"
|
||||||
)
|
)
|
||||||
|
|
||||||
func NewUserDeviceFromConfig(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, routines int) (Device, error) {
|
func NewUserDeviceFromConfig(c *config.C, l *slog.Logger, vpnNetworks []netip.Prefix, routines int) (Device, error) {
|
||||||
@@ -23,11 +25,13 @@ func NewUserDevice(vpnNetworks []netip.Prefix) (Device, error) {
|
|||||||
outboundWriter: ow,
|
outboundWriter: ow,
|
||||||
inboundReader: ir,
|
inboundReader: ir,
|
||||||
inboundWriter: iw,
|
inboundWriter: iw,
|
||||||
|
numReaders: 1,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type UserDevice struct {
|
type UserDevice struct {
|
||||||
vpnNetworks []netip.Prefix
|
vpnNetworks []netip.Prefix
|
||||||
|
numReaders int
|
||||||
|
|
||||||
outboundReader *io.PipeReader
|
outboundReader *io.PipeReader
|
||||||
outboundWriter *io.PipeWriter
|
outboundWriter *io.PipeWriter
|
||||||
@@ -36,6 +40,23 @@ type UserDevice struct {
|
|||||||
inboundWriter *io.PipeWriter
|
inboundWriter *io.PipeWriter
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) Capabilities() tio.Capabilities {
|
||||||
|
return tio.Capabilities{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) Read(p []wire.TunPacket, mem []byte) (int, error) {
|
||||||
|
if len(p) == 0 || len(mem) == 0 {
|
||||||
|
return 0, nil //todo should this be an err?
|
||||||
|
}
|
||||||
|
p[0].Meta = struct{}{}
|
||||||
|
n, err := d.outboundReader.Read(mem)
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p[0].Bytes = mem[:n]
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Activate() error {
|
func (d *UserDevice) Activate() error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -50,20 +71,27 @@ func (d *UserDevice) SupportsMultiqueue() bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) NewMultiQueueReader() (io.ReadWriteCloser, error) {
|
func (d *UserDevice) NewMultiQueueReader() error {
|
||||||
return d, nil
|
d.numReaders++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *UserDevice) Readers() []tio.Queue {
|
||||||
|
out := make([]tio.Queue, d.numReaders)
|
||||||
|
for i := range d.numReaders {
|
||||||
|
out[i] = d
|
||||||
|
}
|
||||||
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
func (d *UserDevice) Pipe() (*io.PipeReader, *io.PipeWriter) {
|
||||||
return d.inboundReader, d.outboundWriter
|
return d.inboundReader, d.outboundWriter
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Read(p []byte) (n int, err error) {
|
|
||||||
return d.outboundReader.Read(p)
|
|
||||||
}
|
|
||||||
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
func (d *UserDevice) Write(p []byte) (n int, err error) {
|
||||||
return d.inboundWriter.Write(p)
|
return d.inboundWriter.Write(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d *UserDevice) Close() error {
|
func (d *UserDevice) Close() error {
|
||||||
d.inboundWriter.Close()
|
d.inboundWriter.Close()
|
||||||
d.outboundWriter.Close()
|
d.outboundWriter.Close()
|
||||||
|
|||||||
@@ -99,12 +99,10 @@ func (p *PKI) reloadCerts(c *config.C, initial bool) *util.ContextualError {
|
|||||||
var currentState *CertState
|
var currentState *CertState
|
||||||
if initial {
|
if initial {
|
||||||
cipher = c.GetString("cipher", "aes")
|
cipher = c.GetString("cipher", "aes")
|
||||||
//TODO: this sucks and we should make it not a global
|
|
||||||
switch cipher {
|
switch cipher {
|
||||||
case "aes":
|
case "aes", "chachapoly":
|
||||||
noiseEndianness = binary.BigEndian
|
// Each post-handshake CipherState in noiseutil hardcodes its own
|
||||||
case "chachapoly":
|
// nonce endianness now, so there's nothing to set up here.
|
||||||
noiseEndianness = binary.LittleEndian
|
|
||||||
default:
|
default:
|
||||||
return util.NewContextualError(
|
return util.NewContextualError(
|
||||||
"unknown cipher",
|
"unknown cipher",
|
||||||
|
|||||||
@@ -1,24 +1,70 @@
|
|||||||
package nebula
|
package nebula
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"net/netip"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/rcrowley/go-metrics"
|
||||||
"github.com/slackhq/nebula/config"
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/header"
|
||||||
|
"github.com/slackhq/nebula/udp"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// holepunchQueueSize buffers the channel that pending holepunchJobs land on after their delay timer fires.
|
||||||
|
const holepunchQueueSize = 64
|
||||||
|
|
||||||
|
// holepunchJob is one scheduled item delivered to the worker goroutine.
|
||||||
|
// - target valid -> send a UDP punch to target. vpnAddr, if set, is the peer's vpn addr carried for log context.
|
||||||
|
// - target invalid, vpnAddr valid -> send an encrypted test packet to vpnAddr (a "punchback").
|
||||||
|
type holepunchJob struct {
|
||||||
|
target netip.AddrPort
|
||||||
|
vpnAddr netip.Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
// lighthouseChecker is the slice of LightHouse that Punchy actually needs.
|
||||||
|
// Defined here so Punchy doesn't take a *LightHouse dependency (LightHouse
|
||||||
|
// already holds a *Punchy, and the bidirectional pointer reference is awkward
|
||||||
|
// even within the same package). Tests can also substitute a fake.
|
||||||
|
type lighthouseChecker interface {
|
||||||
|
IsAnyLighthouseAddr(vpnAddrs []netip.Addr) bool
|
||||||
|
}
|
||||||
|
|
||||||
type Punchy struct {
|
type Punchy struct {
|
||||||
punch atomic.Bool
|
punch atomic.Bool
|
||||||
respond atomic.Bool
|
respond atomic.Bool
|
||||||
delay atomic.Int64
|
delay atomic.Int64
|
||||||
respondDelay atomic.Int64
|
respondDelay atomic.Int64
|
||||||
punchEverything atomic.Bool
|
punchEverything atomic.Bool
|
||||||
l *slog.Logger
|
|
||||||
|
sched *Scheduler[holepunchJob]
|
||||||
|
punchConn udp.Conn
|
||||||
|
metricHolepunchTx metrics.Counter
|
||||||
|
metricPunchyTx metrics.Counter
|
||||||
|
|
||||||
|
ctx context.Context
|
||||||
|
ifce EncWriter
|
||||||
|
hm *HostMap
|
||||||
|
lh lighthouseChecker
|
||||||
|
|
||||||
|
l *slog.Logger
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewPunchyFromConfig(l *slog.Logger, c *config.C) *Punchy {
|
func NewPunchyFromConfig(l *slog.Logger, c *config.C, punchConn udp.Conn) *Punchy {
|
||||||
p := &Punchy{l: l}
|
p := &Punchy{
|
||||||
|
l: l,
|
||||||
|
punchConn: punchConn,
|
||||||
|
sched: NewScheduler[holepunchJob](holepunchQueueSize),
|
||||||
|
metricPunchyTx: metrics.GetOrRegisterCounter("messages.tx.punchy", nil),
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.GetBool("stats.lighthouse_metrics", false) {
|
||||||
|
p.metricHolepunchTx = metrics.GetOrRegisterCounter("messages.tx.holepunch", nil)
|
||||||
|
} else {
|
||||||
|
p.metricHolepunchTx = metrics.NilCounter{}
|
||||||
|
}
|
||||||
|
|
||||||
p.reload(c, true)
|
p.reload(c, true)
|
||||||
c.RegisterReloadCallback(func(c *config.C) {
|
c.RegisterReloadCallback(func(c *config.C) {
|
||||||
@@ -29,7 +75,7 @@ func NewPunchyFromConfig(l *slog.Logger, c *config.C) *Punchy {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (p *Punchy) reload(c *config.C, initial bool) {
|
func (p *Punchy) reload(c *config.C, initial bool) {
|
||||||
if initial {
|
if initial || c.HasChanged("punchy.punch") || c.HasChanged("punchy") {
|
||||||
var yes bool
|
var yes bool
|
||||||
if c.IsSet("punchy.punch") {
|
if c.IsSet("punchy.punch") {
|
||||||
yes = c.GetBool("punchy.punch", false)
|
yes = c.GetBool("punchy.punch", false)
|
||||||
@@ -38,16 +84,15 @@ func (p *Punchy) reload(c *config.C, initial bool) {
|
|||||||
yes = c.GetBool("punchy", false)
|
yes = c.GetBool("punchy", false)
|
||||||
}
|
}
|
||||||
|
|
||||||
p.punch.Store(yes)
|
old := p.punch.Swap(yes)
|
||||||
if yes {
|
switch {
|
||||||
|
case initial && yes:
|
||||||
p.l.Info("punchy enabled")
|
p.l.Info("punchy enabled")
|
||||||
} else {
|
case initial:
|
||||||
p.l.Info("punchy disabled")
|
p.l.Info("punchy disabled")
|
||||||
|
case old != yes:
|
||||||
|
p.l.Info("punchy.punch changed", "punch", yes)
|
||||||
}
|
}
|
||||||
|
|
||||||
} else if c.HasChanged("punchy.punch") || c.HasChanged("punchy") {
|
|
||||||
//TODO: it should be relatively easy to support this, just need to be able to cancel the goroutine and boot it up from here
|
|
||||||
p.l.Warn("Changing punchy.punch with reload is not supported, ignoring.")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if initial || c.HasChanged("punchy.respond") || c.HasChanged("punch_back") {
|
if initial || c.HasChanged("punchy.respond") || c.HasChanged("punch_back") {
|
||||||
@@ -59,52 +104,132 @@ func (p *Punchy) reload(c *config.C, initial bool) {
|
|||||||
yes = c.GetBool("punch_back", false)
|
yes = c.GetBool("punch_back", false)
|
||||||
}
|
}
|
||||||
|
|
||||||
p.respond.Store(yes)
|
old := p.respond.Swap(yes)
|
||||||
|
if !initial && old != yes {
|
||||||
if !initial {
|
p.l.Info("punchy.respond changed", "respond", yes)
|
||||||
p.l.Info("punchy.respond changed", "respond", p.GetRespond())
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
//NOTE: this will not apply to any in progress operations, only the next one
|
//NOTE: this will not apply to any in progress operations, only the next one
|
||||||
if initial || c.HasChanged("punchy.delay") {
|
if initial || c.HasChanged("punchy.delay") {
|
||||||
p.delay.Store((int64)(c.GetDuration("punchy.delay", time.Second)))
|
newDelay := int64(c.GetDuration("punchy.delay", time.Second))
|
||||||
if !initial {
|
old := p.delay.Swap(newDelay)
|
||||||
p.l.Info("punchy.delay changed", "delay", p.GetDelay())
|
if !initial && old != newDelay {
|
||||||
|
p.l.Info("punchy.delay changed", "delay", time.Duration(newDelay))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if initial || c.HasChanged("punchy.target_all_remotes") {
|
if initial || c.HasChanged("punchy.target_all_remotes") {
|
||||||
p.punchEverything.Store(c.GetBool("punchy.target_all_remotes", false))
|
yes := c.GetBool("punchy.target_all_remotes", false)
|
||||||
if !initial {
|
old := p.punchEverything.Swap(yes)
|
||||||
p.l.Info("punchy.target_all_remotes changed", "target_all_remotes", p.GetTargetEverything())
|
if !initial && old != yes {
|
||||||
|
p.l.Info("punchy.target_all_remotes changed", "target_all_remotes", yes)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if initial || c.HasChanged("punchy.respond_delay") {
|
if initial || c.HasChanged("punchy.respond_delay") {
|
||||||
p.respondDelay.Store((int64)(c.GetDuration("punchy.respond_delay", 5*time.Second)))
|
newDelay := int64(c.GetDuration("punchy.respond_delay", 5*time.Second))
|
||||||
if !initial {
|
old := p.respondDelay.Swap(newDelay)
|
||||||
p.l.Info("punchy.respond_delay changed", "respond_delay", p.GetRespondDelay())
|
if !initial && old != newDelay {
|
||||||
|
p.l.Info("punchy.respond_delay changed", "respond_delay", time.Duration(newDelay))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Punchy) GetPunch() bool {
|
// Schedule queues a punch packet to target, to be sent after the configured delay.
|
||||||
return p.punch.Load()
|
// vpnAddr is the peer's vpn addr, used for log context when the packet actually fires.
|
||||||
|
// No-op if target is not a valid AddrPort or if Start has not yet been called. Safe to call from any goroutine.
|
||||||
|
func (p *Punchy) Schedule(target netip.AddrPort, vpnAddr netip.Addr) {
|
||||||
|
if !target.IsValid() || p.ctx == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
p.scheduleJob(holepunchJob{target: target, vpnAddr: vpnAddr}, time.Duration(p.delay.Load()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Punchy) GetRespond() bool {
|
// ScheduleRespond queues a punchback test packet to vpnAddr after the configured respond delay,
|
||||||
return p.respond.Load()
|
// gated on punchy.respond. No-op when respond is disabled or before Start has been called.
|
||||||
|
func (p *Punchy) ScheduleRespond(vpnAddr netip.Addr) {
|
||||||
|
if !p.respond.Load() || p.ctx == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
p.scheduleJob(holepunchJob{vpnAddr: vpnAddr}, time.Duration(p.respondDelay.Load()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Punchy) GetDelay() time.Duration {
|
// scheduleJob delegates to the pooled Scheduler.
|
||||||
return (time.Duration)(p.delay.Load())
|
// The callback observes p.ctx so a job that becomes due after Stop is dropped instead of queued.
|
||||||
|
func (p *Punchy) scheduleJob(job holepunchJob, delay time.Duration) {
|
||||||
|
p.sched.Schedule(p.ctx, job, delay)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Punchy) GetRespondDelay() time.Duration {
|
// SendPunch sends an immediate keepalive punch for an idle hostinfo.
|
||||||
return (time.Duration)(p.respondDelay.Load())
|
// The configured punchy.target_all_remotes mode picks the targets. Gated on punchy.punch and the lighthouse-skip rule
|
||||||
|
// (lighthouses don't get keepalive punches because the regular update interval keeps their NAT state warm).
|
||||||
|
func (p *Punchy) SendPunch(hostinfo *HostInfo) {
|
||||||
|
if !p.punch.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if p.lh.IsAnyLighthouseAddr(hostinfo.vpnAddrs) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.punchEverything.Load() {
|
||||||
|
p.sendPunchToAllRemotes(hostinfo)
|
||||||
|
} else if hostinfo.remote.IsValid() {
|
||||||
|
p.metricPunchyTx.Inc(1)
|
||||||
|
p.punchConn.WriteTo([]byte{1}, hostinfo.remote)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *Punchy) GetTargetEverything() bool {
|
// SendPunchToAll punches every known remote for hostinfo, but only when punchy.target_all_remotes is enabled.
|
||||||
return p.punchEverything.Load()
|
// The connection manager calls this during outbound-only traffic: the outbound traffic itself keeps the primary's
|
||||||
|
// NAT state warm, but non-primary remotes need separate refresh, so we fan out to all of them (the redundant
|
||||||
|
// primary punch is harmless). Gated on punchy.punch and the lighthouse-skip rule.
|
||||||
|
func (p *Punchy) SendPunchToAll(hostinfo *HostInfo) {
|
||||||
|
if !p.punchEverything.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !p.punch.Load() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if p.lh.IsAnyLighthouseAddr(hostinfo.vpnAddrs) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
p.sendPunchToAllRemotes(hostinfo)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Punchy) sendPunchToAllRemotes(hostinfo *HostInfo) {
|
||||||
|
hostinfo.remotes.ForEach(p.hm.GetPreferredRanges(), func(addr netip.AddrPort, preferred bool) {
|
||||||
|
p.metricPunchyTx.Inc(1)
|
||||||
|
p.punchConn.WriteTo([]byte{1}, addr)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Start wires the runtime dependencies and spawns the scheduler worker.
|
||||||
|
func (p *Punchy) Start(ctx context.Context, ifce EncWriter, hm *HostMap, lh lighthouseChecker) {
|
||||||
|
p.ctx = ctx
|
||||||
|
p.ifce = ifce
|
||||||
|
p.hm = hm
|
||||||
|
p.lh = lh
|
||||||
|
|
||||||
|
nb := make([]byte, 12, 12)
|
||||||
|
out := make([]byte, mtu)
|
||||||
|
empty := []byte{0}
|
||||||
|
|
||||||
|
go p.sched.Run(ctx, func(job holepunchJob) {
|
||||||
|
switch {
|
||||||
|
case job.target.IsValid():
|
||||||
|
if p.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
p.l.Debug("Punching", "target", job.target, "vpnAddr", job.vpnAddr)
|
||||||
|
}
|
||||||
|
p.metricHolepunchTx.Inc(1)
|
||||||
|
p.punchConn.WriteTo(empty, job.target)
|
||||||
|
case job.vpnAddr.IsValid():
|
||||||
|
// A nebula test packet to the host trying to contact us.
|
||||||
|
// In the case of a double nat or other difficult scenario, this may help establish a tunnel.
|
||||||
|
if p.l.Enabled(context.Background(), slog.LevelDebug) {
|
||||||
|
p.l.Debug("Sending a nebula test packet", "vpnAddr", job.vpnAddr)
|
||||||
|
}
|
||||||
|
p.ifce.SendMessageToVpnAddr(header.Test, header.TestRequest, job.vpnAddr, []byte(""), nb, out)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
+40
-41
@@ -17,42 +17,42 @@ func TestNewPunchyFromConfig(t *testing.T) {
|
|||||||
c := config.NewC(l)
|
c := config.NewC(l)
|
||||||
|
|
||||||
// Test defaults
|
// Test defaults
|
||||||
p := NewPunchyFromConfig(test.NewLogger(), c)
|
p := NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.False(t, p.GetPunch())
|
assert.False(t, p.punch.Load())
|
||||||
assert.False(t, p.GetRespond())
|
assert.False(t, p.respond.Load())
|
||||||
assert.Equal(t, time.Second, p.GetDelay())
|
assert.Equal(t, time.Second, time.Duration(p.delay.Load()))
|
||||||
assert.Equal(t, 5*time.Second, p.GetRespondDelay())
|
assert.Equal(t, 5*time.Second, time.Duration(p.respondDelay.Load()))
|
||||||
|
|
||||||
// punchy deprecation
|
// punchy deprecation
|
||||||
c.Settings["punchy"] = true
|
c.Settings["punchy"] = true
|
||||||
p = NewPunchyFromConfig(test.NewLogger(), c)
|
p = NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.True(t, p.GetPunch())
|
assert.True(t, p.punch.Load())
|
||||||
|
|
||||||
// punchy.punch
|
// punchy.punch
|
||||||
c.Settings["punchy"] = map[string]any{"punch": true}
|
c.Settings["punchy"] = map[string]any{"punch": true}
|
||||||
p = NewPunchyFromConfig(test.NewLogger(), c)
|
p = NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.True(t, p.GetPunch())
|
assert.True(t, p.punch.Load())
|
||||||
|
|
||||||
// punch_back deprecation
|
// punch_back deprecation
|
||||||
c.Settings["punch_back"] = true
|
c.Settings["punch_back"] = true
|
||||||
p = NewPunchyFromConfig(test.NewLogger(), c)
|
p = NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.True(t, p.GetRespond())
|
assert.True(t, p.respond.Load())
|
||||||
|
|
||||||
// punchy.respond
|
// punchy.respond
|
||||||
c.Settings["punchy"] = map[string]any{"respond": true}
|
c.Settings["punchy"] = map[string]any{"respond": true}
|
||||||
c.Settings["punch_back"] = false
|
c.Settings["punch_back"] = false
|
||||||
p = NewPunchyFromConfig(test.NewLogger(), c)
|
p = NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.True(t, p.GetRespond())
|
assert.True(t, p.respond.Load())
|
||||||
|
|
||||||
// punchy.delay
|
// punchy.delay
|
||||||
c.Settings["punchy"] = map[string]any{"delay": "1m"}
|
c.Settings["punchy"] = map[string]any{"delay": "1m"}
|
||||||
p = NewPunchyFromConfig(test.NewLogger(), c)
|
p = NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.Equal(t, time.Minute, p.GetDelay())
|
assert.Equal(t, time.Minute, time.Duration(p.delay.Load()))
|
||||||
|
|
||||||
// punchy.respond_delay
|
// punchy.respond_delay
|
||||||
c.Settings["punchy"] = map[string]any{"respond_delay": "1m"}
|
c.Settings["punchy"] = map[string]any{"respond_delay": "1m"}
|
||||||
p = NewPunchyFromConfig(test.NewLogger(), c)
|
p = NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.Equal(t, time.Minute, p.GetRespondDelay())
|
assert.Equal(t, time.Minute, time.Duration(p.respondDelay.Load()))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestPunchy_reload(t *testing.T) {
|
func TestPunchy_reload(t *testing.T) {
|
||||||
@@ -61,35 +61,34 @@ func TestPunchy_reload(t *testing.T) {
|
|||||||
delay, _ := time.ParseDuration("1m")
|
delay, _ := time.ParseDuration("1m")
|
||||||
require.NoError(t, c.LoadString(`
|
require.NoError(t, c.LoadString(`
|
||||||
punchy:
|
punchy:
|
||||||
|
punch: false
|
||||||
delay: 1m
|
delay: 1m
|
||||||
respond: false
|
respond: false
|
||||||
`))
|
`))
|
||||||
p := NewPunchyFromConfig(test.NewLogger(), c)
|
p := NewPunchyFromConfig(test.NewLogger(), c, nil)
|
||||||
assert.Equal(t, delay, p.GetDelay())
|
assert.False(t, p.punch.Load())
|
||||||
assert.False(t, p.GetRespond())
|
assert.Equal(t, delay, time.Duration(p.delay.Load()))
|
||||||
|
assert.False(t, p.respond.Load())
|
||||||
|
|
||||||
newDelay, _ := time.ParseDuration("10m")
|
newDelay, _ := time.ParseDuration("10m")
|
||||||
require.NoError(t, c.ReloadConfigString(`
|
require.NoError(t, c.ReloadConfigString(`
|
||||||
punchy:
|
punchy:
|
||||||
|
punch: true
|
||||||
delay: 10m
|
delay: 10m
|
||||||
respond: true
|
respond: true
|
||||||
`))
|
`))
|
||||||
p.reload(c, false)
|
p.reload(c, false)
|
||||||
assert.Equal(t, newDelay, p.GetDelay())
|
assert.True(t, p.punch.Load())
|
||||||
assert.True(t, p.GetRespond())
|
assert.Equal(t, newDelay, time.Duration(p.delay.Load()))
|
||||||
|
assert.True(t, p.respond.Load())
|
||||||
}
|
}
|
||||||
|
|
||||||
// The tests below pin the shape of each log line Punchy produces so changes
|
// The tests below pin the shape of each log line Punchy produces so changes
|
||||||
// cannot silently break whatever operators are grepping for. The assertions
|
// cannot silently break whatever operators are grepping for. The assertions
|
||||||
// are on the structured message + attrs (e.g. "punchy.respond changed" with
|
// are on the structured message + attrs (e.g. "punchy.respond changed" with
|
||||||
// a respond=true field) rather than a formatted string.
|
// a respond=true field) rather than a formatted string. Tests filter by
|
||||||
//
|
// message rather than asserting total entry counts so unrelated info lines
|
||||||
// Punchy.reload also emits a spurious "Changing punchy.punch with reload is
|
// are tolerated without being locked into the format.
|
||||||
// not supported" warning whenever any key under punchy changes, because of
|
|
||||||
// the c.HasChanged("punchy") fallback kept for the deprecated top-level
|
|
||||||
// punchy form. The tests filter by message rather than asserting total
|
|
||||||
// entry counts so that warning is tolerated without being locked into
|
|
||||||
// the format.
|
|
||||||
|
|
||||||
type capturedEntry struct {
|
type capturedEntry struct {
|
||||||
Level slog.Level
|
Level slog.Level
|
||||||
@@ -145,7 +144,7 @@ func TestPunchy_LogFormat_InitialEnabled(t *testing.T) {
|
|||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {punch: true}`))
|
require.NoError(t, c.LoadString(`punchy: {punch: true}`))
|
||||||
|
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
|
|
||||||
entry := findEntry(t, hook.entries, "punchy enabled")
|
entry := findEntry(t, hook.entries, "punchy enabled")
|
||||||
assert.Equal(t, slog.LevelInfo, entry.Level)
|
assert.Equal(t, slog.LevelInfo, entry.Level)
|
||||||
@@ -157,32 +156,32 @@ func TestPunchy_LogFormat_InitialDisabled(t *testing.T) {
|
|||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {punch: false}`))
|
require.NoError(t, c.LoadString(`punchy: {punch: false}`))
|
||||||
|
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
|
|
||||||
entry := findEntry(t, hook.entries, "punchy disabled")
|
entry := findEntry(t, hook.entries, "punchy disabled")
|
||||||
assert.Equal(t, slog.LevelInfo, entry.Level)
|
assert.Equal(t, slog.LevelInfo, entry.Level)
|
||||||
assert.Empty(t, entry.Attrs)
|
assert.Empty(t, entry.Attrs)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestPunchy_LogFormat_ReloadPunchUnsupported(t *testing.T) {
|
func TestPunchy_LogFormat_ReloadPunch(t *testing.T) {
|
||||||
l, hook := newCapturingPunchyLogger(t)
|
l, hook := newCapturingPunchyLogger(t)
|
||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {punch: false}`))
|
require.NoError(t, c.LoadString(`punchy: {punch: false}`))
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
hook.entries = nil
|
hook.entries = nil
|
||||||
|
|
||||||
require.NoError(t, c.ReloadConfigString(`punchy: {punch: true}`))
|
require.NoError(t, c.ReloadConfigString(`punchy: {punch: true}`))
|
||||||
|
|
||||||
entry := findEntry(t, hook.entries, "Changing punchy.punch with reload is not supported, ignoring.")
|
entry := findEntry(t, hook.entries, "punchy.punch changed")
|
||||||
assert.Equal(t, slog.LevelWarn, entry.Level)
|
assert.Equal(t, slog.LevelInfo, entry.Level)
|
||||||
assert.Empty(t, entry.Attrs)
|
assert.Equal(t, map[string]any{"punch": true}, entry.Attrs)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestPunchy_LogFormat_ReloadRespond(t *testing.T) {
|
func TestPunchy_LogFormat_ReloadRespond(t *testing.T) {
|
||||||
l, hook := newCapturingPunchyLogger(t)
|
l, hook := newCapturingPunchyLogger(t)
|
||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {respond: false}`))
|
require.NoError(t, c.LoadString(`punchy: {respond: false}`))
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
hook.entries = nil
|
hook.entries = nil
|
||||||
|
|
||||||
require.NoError(t, c.ReloadConfigString(`punchy: {respond: true}`))
|
require.NoError(t, c.ReloadConfigString(`punchy: {respond: true}`))
|
||||||
@@ -196,7 +195,7 @@ func TestPunchy_LogFormat_ReloadDelay(t *testing.T) {
|
|||||||
l, hook := newCapturingPunchyLogger(t)
|
l, hook := newCapturingPunchyLogger(t)
|
||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {delay: 1s}`))
|
require.NoError(t, c.LoadString(`punchy: {delay: 1s}`))
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
hook.entries = nil
|
hook.entries = nil
|
||||||
|
|
||||||
require.NoError(t, c.ReloadConfigString(`punchy: {delay: 10s}`))
|
require.NoError(t, c.ReloadConfigString(`punchy: {delay: 10s}`))
|
||||||
@@ -210,7 +209,7 @@ func TestPunchy_LogFormat_ReloadTargetAllRemotes(t *testing.T) {
|
|||||||
l, hook := newCapturingPunchyLogger(t)
|
l, hook := newCapturingPunchyLogger(t)
|
||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {target_all_remotes: false}`))
|
require.NoError(t, c.LoadString(`punchy: {target_all_remotes: false}`))
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
hook.entries = nil
|
hook.entries = nil
|
||||||
|
|
||||||
require.NoError(t, c.ReloadConfigString(`punchy: {target_all_remotes: true}`))
|
require.NoError(t, c.ReloadConfigString(`punchy: {target_all_remotes: true}`))
|
||||||
@@ -224,7 +223,7 @@ func TestPunchy_LogFormat_ReloadRespondDelay(t *testing.T) {
|
|||||||
l, hook := newCapturingPunchyLogger(t)
|
l, hook := newCapturingPunchyLogger(t)
|
||||||
c := config.NewC(test.NewLogger())
|
c := config.NewC(test.NewLogger())
|
||||||
require.NoError(t, c.LoadString(`punchy: {respond_delay: 5s}`))
|
require.NoError(t, c.LoadString(`punchy: {respond_delay: 5s}`))
|
||||||
NewPunchyFromConfig(l, c)
|
NewPunchyFromConfig(l, c, nil)
|
||||||
hook.entries = nil
|
hook.entries = nil
|
||||||
|
|
||||||
require.NoError(t, c.ReloadConfigString(`punchy: {respond_delay: 15s}`))
|
require.NoError(t, c.ReloadConfigString(`punchy: {respond_delay: 15s}`))
|
||||||
|
|||||||
@@ -0,0 +1,84 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Scheduler is an allocation-conscious dispatch primitive for delayed work.
|
||||||
|
// Pending items are handed to time.AfterFunc, and ready items land on a worker
|
||||||
|
// channel for centralized dispatch in fire-time order.
|
||||||
|
//
|
||||||
|
// Pick a Scheduler when fire timing matters (exact deadlines, no bucketing) or when the scheduling
|
||||||
|
// rate is uneven enough that idle CPU matters. Each fire is a runtime-spawned goroutine running the callback before
|
||||||
|
// delivering to the worker, which is fine at sparse rates but adds up at line rate.
|
||||||
|
//
|
||||||
|
// Pick a TimerWheel when scheduling is high-rate and uniform: its O(1) insert, internal item cache,
|
||||||
|
// and bucket-batched dispatch are cheaper at scale.
|
||||||
|
// The caller drives the tick loop (Advance/Purge) and pays for fires at bucket boundaries rather than exact deadlines.
|
||||||
|
type Scheduler[T any] struct {
|
||||||
|
queue chan T
|
||||||
|
pool sync.Pool
|
||||||
|
}
|
||||||
|
|
||||||
|
type schedItem[T any] struct {
|
||||||
|
val T
|
||||||
|
ctx context.Context
|
||||||
|
s *Scheduler[T]
|
||||||
|
timer *time.Timer
|
||||||
|
fire func()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewScheduler builds a Scheduler whose worker channel is sized to queueSize.
|
||||||
|
// The buffer absorbs bursts of timers firing close together without
|
||||||
|
// blocking the runtime's callback goroutines on the worker.
|
||||||
|
func NewScheduler[T any](queueSize int) *Scheduler[T] {
|
||||||
|
s := &Scheduler[T]{
|
||||||
|
queue: make(chan T, queueSize),
|
||||||
|
}
|
||||||
|
s.pool.New = func() any {
|
||||||
|
si := &schedItem[T]{s: s}
|
||||||
|
// fire is allocated exactly once per pool-resident item.
|
||||||
|
// The closure captures only `si`, which stays stable for the item's lifetime.
|
||||||
|
si.fire = func() {
|
||||||
|
select {
|
||||||
|
case si.s.queue <- si.val:
|
||||||
|
case <-si.ctx.Done():
|
||||||
|
}
|
||||||
|
var zero T
|
||||||
|
si.val = zero
|
||||||
|
si.ctx = nil
|
||||||
|
si.s.pool.Put(si)
|
||||||
|
}
|
||||||
|
return si
|
||||||
|
}
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
// Schedule arranges item to be delivered to the worker after delay.
|
||||||
|
// The runtime's timer heap handles the wait, so the scheduler itself burns no CPU while idle.
|
||||||
|
// The callback observes ctx: if ctx is cancelled before the timer fires, the item is dropped instead of queued.
|
||||||
|
func (s *Scheduler[T]) Schedule(ctx context.Context, item T, delay time.Duration) {
|
||||||
|
si := s.pool.Get().(*schedItem[T])
|
||||||
|
si.val = item
|
||||||
|
si.ctx = ctx
|
||||||
|
if si.timer == nil {
|
||||||
|
si.timer = time.AfterFunc(delay, si.fire)
|
||||||
|
} else {
|
||||||
|
si.timer.Reset(delay)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run drains the worker queue, calling fn for each item. Returns when ctx is cancelled.
|
||||||
|
// Tests that want deterministic timing should drive the queue directly rather than going through Schedule + Run.
|
||||||
|
func (s *Scheduler[T]) Run(ctx context.Context, fn func(T)) {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case item := <-s.queue:
|
||||||
|
fn(item)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
package nebula
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestScheduler_PooledReuse(t *testing.T) {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
s := NewScheduler[int](16)
|
||||||
|
delivered := make(chan int, 256)
|
||||||
|
go s.Run(ctx, func(item int) { delivered <- item })
|
||||||
|
|
||||||
|
const N = 100
|
||||||
|
for i := 0; i < N; i++ {
|
||||||
|
s.Schedule(ctx, i, time.Millisecond)
|
||||||
|
}
|
||||||
|
|
||||||
|
deadline := time.After(2 * time.Second)
|
||||||
|
got := 0
|
||||||
|
for got < N {
|
||||||
|
select {
|
||||||
|
case <-delivered:
|
||||||
|
got++
|
||||||
|
case <-deadline:
|
||||||
|
t.Fatalf("only %d/%d items delivered", got, N)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkScheduler_Schedule reports allocations per Schedule call.
|
||||||
|
// In steady state the Scheduler's sync.Pool means we should see zero allocs per op once the pool warms up.
|
||||||
|
func BenchmarkScheduler_Schedule(b *testing.B) {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
s := NewScheduler[int](b.N)
|
||||||
|
go s.Run(ctx, func(int) {})
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
s.Schedule(ctx, i, time.Microsecond)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkBareAfterFunc is the comparison baseline.
|
||||||
|
// What we'd pay per Schedule if Punchy called time.AfterFunc directly without the pooled Scheduler.
|
||||||
|
// Allocates a *time.Timer plus a closure each call.
|
||||||
|
func BenchmarkBareAfterFunc(b *testing.B) {
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
queue := make(chan int, b.N)
|
||||||
|
go func() {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-queue:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
i := i
|
||||||
|
time.AfterFunc(time.Microsecond, func() {
|
||||||
|
select {
|
||||||
|
case queue <- i:
|
||||||
|
case <-ctx.Done():
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+38
-21
@@ -27,21 +27,20 @@ type SSHServer struct {
|
|||||||
commands *radix.Tree
|
commands *radix.Tree
|
||||||
listener net.Listener
|
listener net.Listener
|
||||||
|
|
||||||
// Call the cancel() function to stop all active sessions
|
// ctx parents per-Run contexts. Cancelling it (e.g. via Control.Stop) tears the server down even
|
||||||
ctx context.Context
|
// across reloads, since each Run derives a fresh child rather than reusing this one directly.
|
||||||
cancel func()
|
ctx context.Context
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewSSHServer creates a new ssh server rigged with default commands and prepares to listen
|
// NewSSHServer creates a new ssh server rigged with default commands and prepares to listen.
|
||||||
func NewSSHServer(l *slog.Logger) (*SSHServer, error) {
|
// The ssh server's context is parented off the supplied ctx so cancelling it
|
||||||
|
// (e.g. on Control.Stop) tears down active sessions and closes the listener.
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
func NewSSHServer(ctx context.Context, l *slog.Logger) (*SSHServer, error) {
|
||||||
s := &SSHServer{
|
s := &SSHServer{
|
||||||
trustedKeys: make(map[string]map[string]bool),
|
trustedKeys: make(map[string]map[string]bool),
|
||||||
l: l,
|
l: l,
|
||||||
commands: radix.New(),
|
commands: radix.New(),
|
||||||
ctx: ctx,
|
ctx: ctx,
|
||||||
cancel: cancel,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
cc := ssh.CertChecker{
|
cc := ssh.CertChecker{
|
||||||
@@ -151,28 +150,51 @@ func (s *SSHServer) RegisterCommand(c *Command) {
|
|||||||
s.commands.Insert(c.Name, c)
|
s.commands.Insert(c.Name, c)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run begins listening and accepting connections
|
// Run begins listening and accepting connections. Each invocation derives a fresh per-Run context
|
||||||
|
// from the constructor-supplied ctx so a Stop+Run sequence (used by config reload) starts clean
|
||||||
|
// rather than carrying a permanently-cancelled context across runs.
|
||||||
func (s *SSHServer) Run(addr string) error {
|
func (s *SSHServer) Run(addr string) error {
|
||||||
var err error
|
if s.ctx.Err() != nil {
|
||||||
s.listener, err = net.Listen("tcp", addr)
|
return s.ctx.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
listener, err := net.Listen("tcp", addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
// s.listener is the public handle Stop uses to interrupt the active run; listener (the local) is what
|
||||||
|
// this run owns. They start equal but a fast reload may overwrite s.listener with the next run's
|
||||||
|
// listener before this run's watcher fires, so each run must close its own listener via the local
|
||||||
|
// reference.
|
||||||
|
s.listener = listener
|
||||||
|
|
||||||
|
runCtx, cancel := context.WithCancel(s.ctx)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
// Close the listener when this run's context is cancelled. That can come from the parent
|
||||||
|
// (Control.Stop), from Run returning normally (defer cancel above), or transitively when a sibling
|
||||||
|
// run cancels through Stop closing the listener. net.Listener.Close is idempotent so a duplicate
|
||||||
|
// close from Stop is benign.
|
||||||
|
go func() {
|
||||||
|
<-runCtx.Done()
|
||||||
|
if err := listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||||
|
s.l.Warn("Failed to close the sshd listener", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
s.l.Info("SSH server is listening", "sshListener", addr)
|
s.l.Info("SSH server is listening", "sshListener", addr)
|
||||||
|
|
||||||
// Run loops until there is an error
|
// Run loops until there is an error
|
||||||
s.run()
|
s.run(runCtx, listener)
|
||||||
s.closeSessions()
|
|
||||||
|
|
||||||
s.l.Info("SSH server stopped listening")
|
s.l.Info("SSH server stopped listening")
|
||||||
// We don't return an error because run logs for us
|
// We don't return an error because run logs for us
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *SSHServer) run() {
|
func (s *SSHServer) run(ctx context.Context, listener net.Listener) {
|
||||||
for {
|
for {
|
||||||
c, err := s.listener.Accept()
|
c, err := listener.Accept()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if !errors.Is(err, net.ErrClosed) {
|
if !errors.Is(err, net.ErrClosed) {
|
||||||
s.l.Warn("Error in listener, shutting down", "error", err)
|
s.l.Warn("Error in listener, shutting down", "error", err)
|
||||||
@@ -184,7 +206,7 @@ func (s *SSHServer) run() {
|
|||||||
// Ensure that a bad client doesn't hurt us by checking for the parent context
|
// Ensure that a bad client doesn't hurt us by checking for the parent context
|
||||||
// cancellation before calling NewServerConn, and forcing the socket to close when
|
// cancellation before calling NewServerConn, and forcing the socket to close when
|
||||||
// the context is cancelled.
|
// the context is cancelled.
|
||||||
sessionContext, sessionCancel := context.WithCancel(s.ctx)
|
sessionContext, sessionCancel := context.WithCancel(ctx)
|
||||||
go func() {
|
go func() {
|
||||||
<-sessionContext.Done()
|
<-sessionContext.Done()
|
||||||
c.Close()
|
c.Close()
|
||||||
@@ -227,14 +249,9 @@ func (s *SSHServer) run() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *SSHServer) Stop() {
|
func (s *SSHServer) Stop() {
|
||||||
// Close the listener, this will cause all session to terminate as well, see SSHServer.Run
|
|
||||||
if s.listener != nil {
|
if s.listener != nil {
|
||||||
if err := s.listener.Close(); err != nil {
|
if err := s.listener.Close(); err != nil {
|
||||||
s.l.Warn("Failed to close the sshd listener", "error", err)
|
s.l.Warn("Failed to close the sshd listener", "error", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *SSHServer) closeSessions() {
|
|
||||||
s.cancel()
|
|
||||||
}
|
|
||||||
|
|||||||
+17
@@ -8,6 +8,23 @@ import (
|
|||||||
// How many timer objects should be cached
|
// How many timer objects should be cached
|
||||||
const timerCacheMax = 50000
|
const timerCacheMax = 50000
|
||||||
|
|
||||||
|
// TimerWheel is a hashed timing wheel: a fixed slot array indexed by (now + delay) % wheelLen,
|
||||||
|
// with each slot a singly linked list of items due in that bucket.
|
||||||
|
// Adds are O(1), Purges return items in arrival-within-slot order, and an internal cache of TimeoutItems
|
||||||
|
// keeps steady-state inserts allocation-free.
|
||||||
|
//
|
||||||
|
// The TimerWheel does not handle concurrency or lifecycle on its own.
|
||||||
|
// Callers drive Advance/Purge from their own ticker loop, take their own locks (or use LockingTimerWheel),
|
||||||
|
// and decide whether to keep ticking when the wheel is empty.
|
||||||
|
//
|
||||||
|
// Pick a TimerWheel when scheduling is high-rate and uniform: line-rate conntrack inserts,
|
||||||
|
// per-tunnel traffic checks at fixed intervals. O(1) insert plus the item cache means the hot path doesn't allocate.
|
||||||
|
// Items added in the same tick are dispatched together when that slot rotates current,
|
||||||
|
// which amortizes the cost of waking the worker.
|
||||||
|
//
|
||||||
|
// Pick a Scheduler when delay precision matters or scheduling is sparse or uneven.
|
||||||
|
// The wheel rounds requested timeouts up to its tick resolution and clamps anything beyond its wheel duration;
|
||||||
|
// both are silent in this implementation.
|
||||||
type TimerWheel[T any] struct {
|
type TimerWheel[T any] struct {
|
||||||
// Current tick
|
// Current tick
|
||||||
current int
|
current int
|
||||||
|
|||||||
+38
-2
@@ -8,16 +8,49 @@ import (
|
|||||||
|
|
||||||
const MTU = 9001
|
const MTU = 9001
|
||||||
|
|
||||||
|
// MaxWriteBatch is the largest batch any Conn.WriteBatch implementation is
|
||||||
|
// required to accept. Callers SHOULD NOT pass more than this per call; Linux
|
||||||
|
// backends preallocate sendmmsg scratch sized to this value, so exceeding it
|
||||||
|
// only costs additional sendmmsg chunks within a single WriteBatch call.
|
||||||
|
const MaxWriteBatch = 128
|
||||||
|
|
||||||
|
// RxMeta carries per-packet metadata extracted from the RX path (ancillary
|
||||||
|
// data, kernel offload state, etc.) and passed to EncReader callbacks.
|
||||||
|
// Backends that do not produce a particular signal leave its zero value.
|
||||||
|
//
|
||||||
|
// OuterECN is the 2-bit IP-level ECN codepoint stamped on the carrier
|
||||||
|
// datagram (extracted from IP_TOS / IPV6_TCLASS cmsg on Linux). Zero
|
||||||
|
// means Not-ECT, which is also the value backends without ECN RX support
|
||||||
|
// supply on every packet.
|
||||||
|
type RxMeta struct {
|
||||||
|
OuterECN byte
|
||||||
|
}
|
||||||
|
|
||||||
type EncReader func(
|
type EncReader func(
|
||||||
addr netip.AddrPort,
|
addr netip.AddrPort,
|
||||||
payload []byte,
|
payload []byte,
|
||||||
|
meta RxMeta,
|
||||||
)
|
)
|
||||||
|
|
||||||
type Conn interface {
|
type Conn interface {
|
||||||
Rebind() error
|
Rebind() error
|
||||||
LocalAddr() (netip.AddrPort, error)
|
LocalAddr() (netip.AddrPort, error)
|
||||||
ListenOut(r EncReader) error
|
// ListenOut invokes r for each received packet. On batch-capable
|
||||||
|
// backends (recvmmsg), flush is called after each batch is fully
|
||||||
|
// delivered — callers use it to flush per-batch accumulators such as
|
||||||
|
// TUN write coalescers. Single-packet backends call flush after each
|
||||||
|
// packet. flush must not be nil.
|
||||||
|
ListenOut(r EncReader, flush func()) error
|
||||||
WriteTo(b []byte, addr netip.AddrPort) error
|
WriteTo(b []byte, addr netip.AddrPort) error
|
||||||
|
// WriteBatch sends a contiguous batch of packets, each with its own
|
||||||
|
// destination. bufs and addrs must have the same length. outerECNs may
|
||||||
|
// be nil (treated as all-zero / Not-ECT); when non-nil it must have the
|
||||||
|
// same length as bufs, and outerECNs[i] is the 2-bit IP-level ECN
|
||||||
|
// codepoint to set on packet i's outer header. Linux uses sendmmsg(2)
|
||||||
|
// for a single syscall and attaches the value as IP_TOS / IPV6_TCLASS
|
||||||
|
// cmsg; other backends ignore it. Returns on the first error; callers
|
||||||
|
// may observe a partial send if some packets went out before the error.
|
||||||
|
WriteBatch(bufs [][]byte, addrs []netip.AddrPort, outerECNs []byte) error
|
||||||
ReloadConfig(c *config.C)
|
ReloadConfig(c *config.C)
|
||||||
SupportsMultipleReaders() bool
|
SupportsMultipleReaders() bool
|
||||||
Close() error
|
Close() error
|
||||||
@@ -31,7 +64,7 @@ func (NoopConn) Rebind() error {
|
|||||||
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
func (NoopConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
return netip.AddrPort{}, nil
|
return netip.AddrPort{}, nil
|
||||||
}
|
}
|
||||||
func (NoopConn) ListenOut(_ EncReader) error {
|
func (NoopConn) ListenOut(_ EncReader, _ func()) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (NoopConn) SupportsMultipleReaders() bool {
|
func (NoopConn) SupportsMultipleReaders() bool {
|
||||||
@@ -40,6 +73,9 @@ func (NoopConn) SupportsMultipleReaders() bool {
|
|||||||
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
func (NoopConn) WriteTo(_ []byte, _ netip.AddrPort) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
func (NoopConn) WriteBatch(_ [][]byte, _ []netip.AddrPort, _ []byte) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
func (NoopConn) ReloadConfig(_ *config.C) {
|
func (NoopConn) ReloadConfig(_ *config.C) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-2
@@ -5,12 +5,11 @@ package udp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+1
-2
@@ -8,12 +8,11 @@ package udp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,57 @@
|
|||||||
|
//go:build (amd64 || arm64) && !e2e_testing
|
||||||
|
// +build amd64 arm64
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"log/slog"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/slackhq/nebula/config"
|
||||||
|
"github.com/slackhq/nebula/wfp"
|
||||||
|
)
|
||||||
|
|
||||||
|
// wrapWithWDFBypass wraps a Conn so that the first ReloadConfig consults listen.windows_bypass_wdf
|
||||||
|
// and installs a WFP PERMIT filter for the listener's bound UDP port. The session is released when Close runs.
|
||||||
|
func wrapWithWDFBypass(l *slog.Logger, conn Conn) Conn {
|
||||||
|
return &bypassConn{Conn: conn, l: l}
|
||||||
|
}
|
||||||
|
|
||||||
|
type bypassConn struct {
|
||||||
|
Conn
|
||||||
|
|
||||||
|
l *slog.Logger
|
||||||
|
installOnce sync.Once
|
||||||
|
session *wfp.Session
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *bypassConn) ReloadConfig(c *config.C) {
|
||||||
|
b.installOnce.Do(func() {
|
||||||
|
if !c.GetBool("listen.windows_bypass_wdf", true) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
addr, err := b.Conn.LocalAddr()
|
||||||
|
if err != nil {
|
||||||
|
b.l.Warn("Failed to query listener port for WFP bypass", "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s, err := wfp.PermitUDPPort(addr.Port())
|
||||||
|
if err != nil {
|
||||||
|
b.l.Warn("Failed to install WFP bypass filters for listener", "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
b.l.Info("Installed WFP filters bypassing Windows Defender Firewall on UDP listener port",
|
||||||
|
"port", addr.Port())
|
||||||
|
b.session = s
|
||||||
|
})
|
||||||
|
b.Conn.ReloadConfig(c)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *bypassConn) Close() error {
|
||||||
|
if b.session != nil {
|
||||||
|
b.session.Close()
|
||||||
|
b.session = nil
|
||||||
|
}
|
||||||
|
return b.Conn.Close()
|
||||||
|
}
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
//go:build !e2e_testing
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
package udp
|
||||||
|
|
||||||
|
import "log/slog"
|
||||||
|
|
||||||
|
// wrapWithWDFBypass is a no-op on windows-386 since we don't currently build for it.
|
||||||
|
func wrapWithWDFBypass(_ *slog.Logger, conn Conn) Conn {
|
||||||
|
return conn
|
||||||
|
}
|
||||||
+12
-2
@@ -140,6 +140,15 @@ func (u *StdConn) WriteTo(b []byte, ap netip.AddrPort) error {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *StdConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -165,7 +174,7 @@ func NewUDPStatsEmitter(udpConns []Conn) func() {
|
|||||||
return func() {}
|
return func() {}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader) error {
|
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
@@ -179,7 +188,8 @@ func (u *StdConn) ListenOut(r EncReader) error {
|
|||||||
u.l.Error("unexpected udp socket receive error", "error", err)
|
u.l.Error("unexpected udp socket receive error", "error", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+12
-2
@@ -44,6 +44,15 @@ func (u *GenericConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *GenericConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if _, err := u.UDPConn.WriteToUDPAddrPort(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *GenericConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
a := u.UDPConn.LocalAddr()
|
a := u.UDPConn.LocalAddr()
|
||||||
|
|
||||||
@@ -73,7 +82,7 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *GenericConn) ListenOut(r EncReader) error {
|
func (u *GenericConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -93,7 +102,8 @@ func (u *GenericConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n])
|
r(netip.AddrPortFrom(rua.Addr().Unmap(), rua.Port()), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+171
-14
@@ -24,6 +24,22 @@ type StdConn struct {
|
|||||||
isV4 bool
|
isV4 bool
|
||||||
l *slog.Logger
|
l *slog.Logger
|
||||||
batch int
|
batch int
|
||||||
|
|
||||||
|
// sendmmsg scratch. Each queue has its own StdConn, so no locking is
|
||||||
|
// needed. Sized to MaxWriteBatch at construction; WriteBatch chunks
|
||||||
|
// larger inputs.
|
||||||
|
writeMsgs []rawMessage
|
||||||
|
writeIovs []iovec
|
||||||
|
writeNames [][]byte
|
||||||
|
|
||||||
|
// sendmmsg(2) callback state. sendmmsgCB is bound once in NewListener
|
||||||
|
// to the sendmmsgRun method value so passing it to rawConn.Write does
|
||||||
|
// not allocate a fresh closure per send; sendmmsgN/Sent/Errno carry
|
||||||
|
// the inputs and outputs across the call without escaping locals.
|
||||||
|
sendmmsgCB func(fd uintptr) bool
|
||||||
|
sendmmsgN int
|
||||||
|
sendmmsgSent int
|
||||||
|
sendmmsgErrno syscall.Errno
|
||||||
}
|
}
|
||||||
|
|
||||||
func setReusePort(network, address string, c syscall.RawConn) error {
|
func setReusePort(network, address string, c syscall.RawConn) error {
|
||||||
@@ -70,9 +86,23 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
|||||||
}
|
}
|
||||||
out.isV4 = af == unix.AF_INET
|
out.isV4 = af == unix.AF_INET
|
||||||
|
|
||||||
|
out.prepareWriteMessages(MaxWriteBatch)
|
||||||
|
out.sendmmsgCB = out.sendmmsgRun
|
||||||
|
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) prepareWriteMessages(n int) {
|
||||||
|
u.writeMsgs = make([]rawMessage, n)
|
||||||
|
u.writeIovs = make([]iovec, n)
|
||||||
|
u.writeNames = make([][]byte, n)
|
||||||
|
|
||||||
|
for i := range u.writeMsgs {
|
||||||
|
u.writeNames[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
u.writeMsgs[i].Hdr.Name = &u.writeNames[i][0]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) SupportsMultipleReaders() bool {
|
func (u *StdConn) SupportsMultipleReaders() bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -171,7 +201,7 @@ func recvmmsg(fd uintptr, msgs []rawMessage) (int, bool, error) {
|
|||||||
return int(n), true, nil
|
return int(n), true, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) listenOutSingle(r EncReader) error {
|
func (u *StdConn) listenOutSingle(r EncReader, flush func()) error {
|
||||||
var err error
|
var err error
|
||||||
var n int
|
var n int
|
||||||
var from netip.AddrPort
|
var from netip.AddrPort
|
||||||
@@ -183,16 +213,33 @@ func (u *StdConn) listenOutSingle(r EncReader) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
from = netip.AddrPortFrom(from.Addr().Unmap(), from.Port())
|
||||||
r(from, buffer[:n])
|
// listenOutSingle uses ReadFromUDPAddrPort which discards cmsgs,
|
||||||
|
// so the outer ECN field is not visible on this path. Zero RxMeta
|
||||||
|
// (Not-ECT) means RFC 6040 combine is a no-op.
|
||||||
|
r(from, buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) listenOutBatch(r EncReader) error {
|
// readSockaddr decodes the source address out of a recvmmsg name buffer
|
||||||
|
func (u *StdConn) readSockaddr(name []byte) netip.AddrPort {
|
||||||
var ip netip.Addr
|
var ip netip.Addr
|
||||||
|
// It's ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
||||||
|
if u.isV4 {
|
||||||
|
ip, _ = netip.AddrFromSlice(name[4:8])
|
||||||
|
} else {
|
||||||
|
ip, _ = netip.AddrFromSlice(name[8:24])
|
||||||
|
}
|
||||||
|
return netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(name[2:4]))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *StdConn) listenOutBatch(r EncReader, flush func()) error {
|
||||||
var n int
|
var n int
|
||||||
var operr error
|
var operr error
|
||||||
|
|
||||||
msgs, buffers, names := u.PrepareRawMessages(u.batch)
|
bufSize := MTU
|
||||||
|
cmsgSpace := 0
|
||||||
|
msgs, buffers, names, _ := u.PrepareRawMessages(u.batch, bufSize, cmsgSpace)
|
||||||
|
|
||||||
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
//reader needs to capture variables from this function, since it's used as a lambda with rawConn.Read
|
||||||
//defining it outside the loop so it gets re-used
|
//defining it outside the loop so it gets re-used
|
||||||
@@ -211,22 +258,18 @@ func (u *StdConn) listenOutBatch(r EncReader) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for i := 0; i < n; i++ {
|
for i := 0; i < n; i++ {
|
||||||
// Its ok to skip the ok check here, the slicing is the only error that can occur and it will panic
|
r(u.readSockaddr(names[i]), buffers[i][:msgs[i].Len], RxMeta{})
|
||||||
if u.isV4 {
|
|
||||||
ip, _ = netip.AddrFromSlice(names[i][4:8])
|
|
||||||
} else {
|
|
||||||
ip, _ = netip.AddrFromSlice(names[i][8:24])
|
|
||||||
}
|
|
||||||
r(netip.AddrPortFrom(ip.Unmap(), binary.BigEndian.Uint16(names[i][2:4])), buffers[i][:msgs[i].Len])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) ListenOut(r EncReader) error {
|
func (u *StdConn) ListenOut(r EncReader, flush func()) error {
|
||||||
if u.batch == 1 {
|
if u.batch == 1 {
|
||||||
return u.listenOutSingle(r)
|
return u.listenOutSingle(r, flush)
|
||||||
} else {
|
} else {
|
||||||
return u.listenOutBatch(r)
|
return u.listenOutBatch(r, flush)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -235,6 +278,120 @@ func (u *StdConn) WriteTo(b []byte, ip netip.AddrPort) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// WriteBatch sends bufs via sendmmsg(2) using the preallocated scratch on
|
||||||
|
// StdConn. If supported, consecutive packets to the same destination with
|
||||||
|
// matching segment sizes (all but possibly the last) are coalesced into a
|
||||||
|
// single mmsghdr entry
|
||||||
|
//
|
||||||
|
// If sendmmsg returns an error and zero entries went out, we fall back to
|
||||||
|
// per-packet WriteTo for that chunk so the caller still gets best-effort
|
||||||
|
// delivery. On a partial send we resume at the first un-acked entry on
|
||||||
|
// the next iteration.
|
||||||
|
func (u *StdConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i := 0; i < len(bufs); {
|
||||||
|
chunk := min(len(bufs)-i, len(u.writeMsgs))
|
||||||
|
|
||||||
|
for k := 0; k < chunk; k++ {
|
||||||
|
u.writeIovs[k].Base = &bufs[i+k][0]
|
||||||
|
setIovLen(&u.writeIovs[k], len(bufs[i+k]))
|
||||||
|
|
||||||
|
nlen, err := writeSockaddr(u.writeNames[k], addrs[i+k], u.isV4)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
hdr := &u.writeMsgs[k].Hdr
|
||||||
|
hdr.Iov = &u.writeIovs[k]
|
||||||
|
setMsgIovlen(hdr, 1)
|
||||||
|
hdr.Namelen = uint32(nlen)
|
||||||
|
}
|
||||||
|
|
||||||
|
sent, serr := u.sendmmsg(chunk)
|
||||||
|
if serr != nil && sent <= 0 {
|
||||||
|
// sendmmsg returns -1 / sent=0 when entry 0 itself failed; log
|
||||||
|
// that entry's destination and fall back to per-packet WriteTo
|
||||||
|
// for the whole chunk so the caller still gets best-effort
|
||||||
|
// delivery without duplicating packets the kernel accepted.
|
||||||
|
u.l.Warn("sendmmsg failed, falling back to per-packet WriteTo",
|
||||||
|
"err", serr,
|
||||||
|
"entries", chunk,
|
||||||
|
"entry0_dst", addrs[i],
|
||||||
|
"isV4", u.isV4,
|
||||||
|
)
|
||||||
|
for k := 0; k < chunk; k++ {
|
||||||
|
if werr := u.WriteTo(bufs[i+k], addrs[i+k]); werr != nil {
|
||||||
|
return werr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
i += chunk
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
i += sent
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendmmsg issues sendmmsg(2) against the first n entries of u.writeMsgs.
|
||||||
|
// The bound u.sendmmsgCB is passed to rawConn.Write so no closure is
|
||||||
|
// allocated per call; inputs and outputs ride on the StdConn fields.
|
||||||
|
func (u *StdConn) sendmmsg(n int) (int, error) {
|
||||||
|
u.sendmmsgN = n
|
||||||
|
u.sendmmsgSent = 0
|
||||||
|
u.sendmmsgErrno = 0
|
||||||
|
if err := u.rawConn.Write(u.sendmmsgCB); err != nil {
|
||||||
|
return u.sendmmsgSent, err
|
||||||
|
}
|
||||||
|
if u.sendmmsgErrno != 0 {
|
||||||
|
return u.sendmmsgSent, &net.OpError{Op: "sendmmsg", Err: u.sendmmsgErrno}
|
||||||
|
}
|
||||||
|
return u.sendmmsgSent, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sendmmsgRun is the rawConn.Write callback. It is bound once into
|
||||||
|
// u.sendmmsgCB at construction so it stays alloc-free in the hot path;
|
||||||
|
// inputs (sendmmsgN) and outputs (sendmmsgSent, sendmmsgErrno) ride on
|
||||||
|
// the receiver rather than escaping locals.
|
||||||
|
func (u *StdConn) sendmmsgRun(fd uintptr) bool {
|
||||||
|
r1, _, errno := unix.Syscall6(unix.SYS_SENDMMSG, fd,
|
||||||
|
uintptr(unsafe.Pointer(&u.writeMsgs[0])), uintptr(u.sendmmsgN),
|
||||||
|
0, 0, 0,
|
||||||
|
)
|
||||||
|
if errno == syscall.EAGAIN || errno == syscall.EWOULDBLOCK {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
u.sendmmsgSent = int(r1)
|
||||||
|
u.sendmmsgErrno = errno
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeSockaddr encodes addr into buf (which must be at least
|
||||||
|
// SizeofSockaddrInet6 bytes). Returns the number of bytes used. If isV4 is
|
||||||
|
// true and addr is not a v4 (or v4-in-v6) address, returns an error.
|
||||||
|
func writeSockaddr(buf []byte, addr netip.AddrPort, isV4 bool) (int, error) {
|
||||||
|
ap := addr.Addr().Unmap()
|
||||||
|
if isV4 {
|
||||||
|
if !ap.Is4() {
|
||||||
|
return 0, ErrInvalidIPv6RemoteForSocket
|
||||||
|
}
|
||||||
|
// struct sockaddr_in: { sa_family_t(2), in_port_t(2, BE), in_addr(4), zero(8) }
|
||||||
|
// sa_family is host endian.
|
||||||
|
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET)
|
||||||
|
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||||
|
ip4 := ap.As4()
|
||||||
|
copy(buf[4:8], ip4[:])
|
||||||
|
clear(buf[8:16])
|
||||||
|
return unix.SizeofSockaddrInet4, nil
|
||||||
|
}
|
||||||
|
// struct sockaddr_in6: { sa_family_t(2), in_port_t(2, BE), flowinfo(4), in6_addr(16), scope_id(4) }
|
||||||
|
binary.NativeEndian.PutUint16(buf[0:2], unix.AF_INET6)
|
||||||
|
binary.BigEndian.PutUint16(buf[2:4], addr.Port())
|
||||||
|
binary.NativeEndian.PutUint32(buf[4:8], 0)
|
||||||
|
ip6 := addr.Addr().As16()
|
||||||
|
copy(buf[8:24], ip6[:])
|
||||||
|
binary.NativeEndian.PutUint32(buf[24:28], 0)
|
||||||
|
return unix.SizeofSockaddrInet6, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *StdConn) ReloadConfig(c *config.C) {
|
func (u *StdConn) ReloadConfig(c *config.C) {
|
||||||
b := c.GetInt("listen.read_buffer", 0)
|
b := c.GetInt("listen.read_buffer", 0)
|
||||||
if b > 0 {
|
if b > 0 {
|
||||||
|
|||||||
+29
-3
@@ -30,13 +30,18 @@ type rawMessage struct {
|
|||||||
Len uint32
|
Len uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||||
msgs := make([]rawMessage, n)
|
msgs := make([]rawMessage, n)
|
||||||
buffers := make([][]byte, n)
|
buffers := make([][]byte, n)
|
||||||
names := make([][]byte, n)
|
names := make([][]byte, n)
|
||||||
|
|
||||||
|
var cmsgs []byte
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
cmsgs = make([]byte, n*cmsgSpace)
|
||||||
|
}
|
||||||
|
|
||||||
for i := range msgs {
|
for i := range msgs {
|
||||||
buffers[i] = make([]byte, MTU)
|
buffers[i] = make([]byte, bufSize)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -48,7 +53,28 @@ func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
|||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
msgs[i].Hdr.Name = &names[i][0]
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||||
|
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||||
|
msgs[i].Hdr.Controllen = uint32(cmsgSpace)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return msgs, buffers, names
|
return msgs, buffers, names, cmsgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func setIovLen(v *iovec, n int) {
|
||||||
|
v.Len = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgIovlen(m *msghdr, n int) {
|
||||||
|
m.Iovlen = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgControllen(m *msghdr, n int) {
|
||||||
|
m.Controllen = uint32(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||||
|
h.Len = uint32(n)
|
||||||
}
|
}
|
||||||
|
|||||||
+29
-3
@@ -33,13 +33,18 @@ type rawMessage struct {
|
|||||||
Pad0 [4]byte
|
Pad0 [4]byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
func (u *StdConn) PrepareRawMessages(n, bufSize, cmsgSpace int) ([]rawMessage, [][]byte, [][]byte, []byte) {
|
||||||
msgs := make([]rawMessage, n)
|
msgs := make([]rawMessage, n)
|
||||||
buffers := make([][]byte, n)
|
buffers := make([][]byte, n)
|
||||||
names := make([][]byte, n)
|
names := make([][]byte, n)
|
||||||
|
|
||||||
|
var cmsgs []byte
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
cmsgs = make([]byte, n*cmsgSpace)
|
||||||
|
}
|
||||||
|
|
||||||
for i := range msgs {
|
for i := range msgs {
|
||||||
buffers[i] = make([]byte, MTU)
|
buffers[i] = make([]byte, bufSize)
|
||||||
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
names[i] = make([]byte, unix.SizeofSockaddrInet6)
|
||||||
|
|
||||||
vs := []iovec{
|
vs := []iovec{
|
||||||
@@ -51,7 +56,28 @@ func (u *StdConn) PrepareRawMessages(n int) ([]rawMessage, [][]byte, [][]byte) {
|
|||||||
|
|
||||||
msgs[i].Hdr.Name = &names[i][0]
|
msgs[i].Hdr.Name = &names[i][0]
|
||||||
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
msgs[i].Hdr.Namelen = uint32(len(names[i]))
|
||||||
|
|
||||||
|
if cmsgSpace > 0 {
|
||||||
|
msgs[i].Hdr.Control = &cmsgs[i*cmsgSpace]
|
||||||
|
msgs[i].Hdr.Controllen = uint64(cmsgSpace)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return msgs, buffers, names
|
return msgs, buffers, names, cmsgs
|
||||||
|
}
|
||||||
|
|
||||||
|
func setIovLen(v *iovec, n int) {
|
||||||
|
v.Len = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgIovlen(m *msghdr, n int) {
|
||||||
|
m.Iovlen = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setMsgControllen(m *msghdr, n int) {
|
||||||
|
m.Controllen = uint64(n)
|
||||||
|
}
|
||||||
|
|
||||||
|
func setCmsgLen(h *unix.Cmsghdr, n int) {
|
||||||
|
h.Len = uint64(n)
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-2
@@ -7,12 +7,11 @@ package udp
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"syscall"
|
"syscall"
|
||||||
|
|
||||||
"log/slog"
|
|
||||||
|
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+12
-2
@@ -140,7 +140,7 @@ func (u *RIOConn) bind(l *slog.Logger, sa windows.Sockaddr) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (u *RIOConn) ListenOut(r EncReader) error {
|
func (u *RIOConn) ListenOut(r EncReader, flush func()) error {
|
||||||
buffer := make([]byte, MTU)
|
buffer := make([]byte, MTU)
|
||||||
|
|
||||||
var lastRecvErr time.Time
|
var lastRecvErr time.Time
|
||||||
@@ -161,7 +161,8 @@ func (u *RIOConn) ListenOut(r EncReader) error {
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n])
|
r(netip.AddrPortFrom(netip.AddrFrom16(rua.Addr).Unmap(), (rua.Port>>8)|((rua.Port&0xff)<<8)), buffer[:n], RxMeta{})
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -316,6 +317,15 @@ func (u *RIOConn) WriteTo(buf []byte, ip netip.AddrPort) error {
|
|||||||
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
return winrio.SendEx(u.rq, dataBuffer, 1, nil, addressBuffer, nil, nil, 0, 0)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (u *RIOConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
func (u *RIOConn) LocalAddr() (netip.AddrPort, error) {
|
||||||
sa, err := windows.Getsockname(u.sock)
|
sa, err := windows.Getsockname(u.sock)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
+11
-2
@@ -157,15 +157,24 @@ func (u *TesterConn) WriteTo(b []byte, addr netip.AddrPort) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
func (u *TesterConn) WriteBatch(bufs [][]byte, addrs []netip.AddrPort, _ []byte) error {
|
||||||
|
for i, b := range bufs {
|
||||||
|
if err := u.WriteTo(b, addrs[i]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (u *TesterConn) ListenOut(r EncReader) error {
|
func (u *TesterConn) ListenOut(r EncReader, flush func()) error {
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-u.done:
|
case <-u.done:
|
||||||
return os.ErrClosed
|
return os.ErrClosed
|
||||||
case p := <-u.RxPackets:
|
case p := <-u.RxPackets:
|
||||||
r(p.From, p.Data)
|
r(p.From, p.Data, RxMeta{})
|
||||||
p.Release()
|
p.Release()
|
||||||
|
flush()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+9
-4
@@ -19,13 +19,18 @@ func NewListener(l *slog.Logger, ip netip.Addr, port int, multi bool, batch int)
|
|||||||
return nil, fmt.Errorf("multiple udp listeners not supported on windows")
|
return nil, fmt.Errorf("multiple udp listeners not supported on windows")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var conn Conn
|
||||||
rc, err := NewRIOListener(l, ip, port)
|
rc, err := NewRIOListener(l, ip, port)
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return rc, nil
|
conn = rc
|
||||||
|
} else {
|
||||||
|
l.Error("Falling back to standard udp sockets", "error", err)
|
||||||
|
conn, err = NewGenericListener(l, ip, port, multi, batch)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
return wrapWithWDFBypass(l, conn), nil
|
||||||
l.Error("Falling back to standard udp sockets", "error", err)
|
|
||||||
return NewGenericListener(l, ip, port, multi, batch)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewListenConfig(multi bool) net.ListenConfig {
|
func NewListenConfig(multi bool) net.ListenConfig {
|
||||||
|
|||||||
@@ -0,0 +1,377 @@
|
|||||||
|
//go:build (amd64 || arm64) && !e2e_testing
|
||||||
|
// +build amd64 arm64
|
||||||
|
// +build !e2e_testing
|
||||||
|
|
||||||
|
// Package wfp installs Windows Filtering Platform (WFP) PERMIT filters in a dynamic, session-scoped sublayer.
|
||||||
|
// Because WFP sits below Windows Defender Firewall, a high-weight permit at FWPM_LAYER_ALE_AUTH_RECV_ACCEPT_V4/V6 lets
|
||||||
|
// the matching inbound traffic through regardless of WDF rules.
|
||||||
|
//
|
||||||
|
// Each Session owns its own engine handle. When the handle closes, every dynamic object added during the session
|
||||||
|
// is auto-deleted by Windows, so there are no orphaned filters.
|
||||||
|
//
|
||||||
|
// Type definitions and constants are derived from the wireguard-windows firewall package (MIT).
|
||||||
|
// Only the subset we exercise is reproduced.
|
||||||
|
package wfp
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
)
|
||||||
|
|
||||||
|
// FWPM layer GUIDs (fwpmu.h).
|
||||||
|
//
|
||||||
|
// FWPM_LAYER_ALE_AUTH_RECV_ACCEPT_V4 = e1cd9fe7-f4b5-4273-96c0-592e487b8650
|
||||||
|
// FWPM_LAYER_ALE_AUTH_RECV_ACCEPT_V6 = a3b42c97-9f04-4672-b87e-cee9c483257f
|
||||||
|
var (
|
||||||
|
fwpmLayerAleAuthRecvAcceptV4 = windows.GUID{
|
||||||
|
Data1: 0xe1cd9fe7, Data2: 0xf4b5, Data3: 0x4273,
|
||||||
|
Data4: [8]byte{0x96, 0xc0, 0x59, 0x2e, 0x48, 0x7b, 0x86, 0x50},
|
||||||
|
}
|
||||||
|
fwpmLayerAleAuthRecvAcceptV6 = windows.GUID{
|
||||||
|
Data1: 0xa3b42c97, Data2: 0x9f04, Data3: 0x4672,
|
||||||
|
Data4: [8]byte{0xb8, 0x7e, 0xce, 0xe9, 0xc4, 0x83, 0x25, 0x7f},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
// FWPM_CONDITION_IP_LOCAL_INTERFACE = 4cd62a49-59c3-4969-b7f3-bda5d32890a4
|
||||||
|
var fwpmConditionIPLocalInterface = windows.GUID{
|
||||||
|
Data1: 0x4cd62a49, Data2: 0x59c3, Data3: 0x4969,
|
||||||
|
Data4: [8]byte{0xb7, 0xf3, 0xbd, 0xa5, 0xd3, 0x28, 0x90, 0xa4},
|
||||||
|
}
|
||||||
|
|
||||||
|
// FWPM_CONDITION_IP_PROTOCOL = 3971ef2b-623e-4f9a-8cb1-6e79b806b9a7
|
||||||
|
var fwpmConditionIPProtocol = windows.GUID{
|
||||||
|
Data1: 0x3971ef2b, Data2: 0x623e, Data3: 0x4f9a,
|
||||||
|
Data4: [8]byte{0x8c, 0xb1, 0x6e, 0x79, 0xb8, 0x06, 0xb9, 0xa7},
|
||||||
|
}
|
||||||
|
|
||||||
|
// FWPM_CONDITION_IP_LOCAL_PORT = 0c1ba1af-5765-453f-af22-a8f791ac775b
|
||||||
|
var fwpmConditionIPLocalPort = windows.GUID{
|
||||||
|
Data1: 0x0c1ba1af, Data2: 0x5765, Data3: 0x453f,
|
||||||
|
Data4: [8]byte{0xaf, 0x22, 0xa8, 0xf7, 0x91, 0xac, 0x77, 0x5b},
|
||||||
|
}
|
||||||
|
|
||||||
|
// IPPROTO_UDP from in.h.
|
||||||
|
const ipprotoUDP uint8 = 17
|
||||||
|
|
||||||
|
// FWP_ACTION_TYPE values (fwptypes.h). PERMIT is terminating.
|
||||||
|
const fwpActionPermit uint32 = 0x00001002 // 0x2 | FWP_ACTION_FLAG_TERMINATING(0x1000)
|
||||||
|
|
||||||
|
// FWP_DATA_TYPE values we use.
|
||||||
|
const (
|
||||||
|
fwpEmpty uint32 = 0
|
||||||
|
fwpUint8 uint32 = 1
|
||||||
|
fwpUint16 uint32 = 2
|
||||||
|
fwpUint64 uint32 = 4
|
||||||
|
)
|
||||||
|
|
||||||
|
// FWP_MATCH_TYPE values.
|
||||||
|
const fwpMatchEqual uint32 = 0
|
||||||
|
|
||||||
|
// FWPM_SESSION flags.
|
||||||
|
const fwpmSessionFlagDynamic uint32 = 0x1
|
||||||
|
|
||||||
|
// FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT prevents lower-priority filters in other sublayers,
|
||||||
|
// notably Windows Defender Firewall's MPSSVC_WF sublayer, which shares our 0xFFFF weight from overriding this PERMIT.
|
||||||
|
// Without it, a default WDF block at the same sublayer weight can still win arbitration.
|
||||||
|
const fwpmFilterFlagClearActionRight uint32 = 0x8
|
||||||
|
|
||||||
|
// RPC authentication.
|
||||||
|
// RPC_C_AUTHN_WINNT works on workgroup machines with no domain context
|
||||||
|
// RPC_C_AUTHN_DEFAULT falls back through a chain that can land on something WFP doesn't accept on a fresh box.
|
||||||
|
const rpcCAuthnWinNT uint32 = 10
|
||||||
|
|
||||||
|
// fwpByteBlob (FWP_BYTE_BLOB). 16 bytes on 64-bit.
|
||||||
|
type fwpByteBlob struct {
|
||||||
|
size uint32
|
||||||
|
_ uint32 // padding
|
||||||
|
data *uint8
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpValue0 / FWP_CONDITION_VALUE0 layout. 16 bytes on 64-bit.
|
||||||
|
// The union is pointer-sized; types <= 32 bits (UINT8/16/32, INT8/16/32, float) live inline in the low bytes
|
||||||
|
// of `value`, while UINT64/INT64/double and aggregate types are stored *by pointer*, even on 64-bit, where the
|
||||||
|
// union member is declared as UINT64*. So when populating an FWP_UINT64 condition, pass
|
||||||
|
// uintptr(unsafe.Pointer(&luidVar)) instead of the LUID inline.
|
||||||
|
type fwpValue0 struct {
|
||||||
|
type_ uint32
|
||||||
|
_ uint32 // padding before union to 8-byte alignment
|
||||||
|
value uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmDisplayData0 / FWPM_DISPLAY_DATA0. 16 bytes on 64-bit.
|
||||||
|
type fwpmDisplayData0 struct {
|
||||||
|
name *uint16
|
||||||
|
description *uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmAction0 / FWPM_ACTION0. 20 bytes; no leading padding because actionType
|
||||||
|
// is uint32 and GUID's first field is uint32.
|
||||||
|
type fwpmAction0 struct {
|
||||||
|
actionType uint32
|
||||||
|
filterType windows.GUID
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmFilterCondition0. 40 bytes on 64-bit.
|
||||||
|
type fwpmFilterCondition0 struct {
|
||||||
|
fieldKey windows.GUID // 16
|
||||||
|
matchType uint32 // 4
|
||||||
|
_ uint32 // 4 padding
|
||||||
|
conditionValue fwpValue0 // 16
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmFilter0. 200 bytes on 64-bit.
|
||||||
|
type fwpmFilter0 struct {
|
||||||
|
filterKey windows.GUID
|
||||||
|
displayData fwpmDisplayData0
|
||||||
|
flags uint32
|
||||||
|
_ uint32 // padding before *GUID
|
||||||
|
providerKey *windows.GUID
|
||||||
|
providerData fwpByteBlob
|
||||||
|
layerKey windows.GUID
|
||||||
|
subLayerKey windows.GUID
|
||||||
|
weight fwpValue0
|
||||||
|
numFilterConditions uint32
|
||||||
|
_ uint32 // padding before pointer
|
||||||
|
filterCondition *fwpmFilterCondition0
|
||||||
|
action fwpmAction0
|
||||||
|
_ [4]byte // layout correction
|
||||||
|
providerContextKey windows.GUID
|
||||||
|
reserved *windows.GUID
|
||||||
|
filterID uint64
|
||||||
|
effectiveWeight fwpValue0
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmSublayer0. 72 bytes on 64-bit.
|
||||||
|
type fwpmSublayer0 struct {
|
||||||
|
subLayerKey windows.GUID
|
||||||
|
displayData fwpmDisplayData0
|
||||||
|
flags uint32
|
||||||
|
_ uint32 // padding before *GUID
|
||||||
|
providerKey *windows.GUID
|
||||||
|
providerData fwpByteBlob
|
||||||
|
weight uint16
|
||||||
|
_ [6]byte // padding to 72 bytes
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmSession0. 72 bytes on 64-bit.
|
||||||
|
type fwpmSession0 struct {
|
||||||
|
sessionKey windows.GUID
|
||||||
|
displayData fwpmDisplayData0
|
||||||
|
flags uint32
|
||||||
|
txnWaitTimeoutInMSec uint32
|
||||||
|
processId uint32
|
||||||
|
_ uint32 // padding before *SID
|
||||||
|
sid *windows.SID
|
||||||
|
username *uint16
|
||||||
|
kernelMode uint8
|
||||||
|
_ [7]byte // tail padding
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpuclnt.dll bindings. Only the calls we use.
|
||||||
|
var (
|
||||||
|
modFwpuclnt = windows.NewLazySystemDLL("fwpuclnt.dll")
|
||||||
|
procFwpmEngineOpen0 = modFwpuclnt.NewProc("FwpmEngineOpen0")
|
||||||
|
procFwpmEngineClose0 = modFwpuclnt.NewProc("FwpmEngineClose0")
|
||||||
|
procFwpmSubLayerAdd0 = modFwpuclnt.NewProc("FwpmSubLayerAdd0")
|
||||||
|
procFwpmFilterAdd0 = modFwpuclnt.NewProc("FwpmFilterAdd0")
|
||||||
|
)
|
||||||
|
|
||||||
|
// Session holds the WFP engine handle for a single bypass operation. The handle owns a dynamic session:
|
||||||
|
// when it is closed, every WFP object added during the session (sublayer + filters) is automatically deleted by
|
||||||
|
// Windows. That gives us correct cleanup even if the host process is killed hard between Permit* and Close.
|
||||||
|
type Session struct {
|
||||||
|
engine uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close releases the engine handle. Windows deletes every dynamic object (sublayer + filters) the session installed.
|
||||||
|
// Safe to call on a nil receiver.
|
||||||
|
func (s *Session) Close() {
|
||||||
|
if s == nil || s.engine == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
procFwpmEngineClose0.Call(s.engine)
|
||||||
|
s.engine = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// PermitInterface installs PERMIT filters at FWPM_LAYER_ALE_AUTH_RECV_ACCEPT_V4 and _V6 scoped to the given network
|
||||||
|
// interface LUID. Inbound traffic on that interface bypasses Windows Defender Firewall.
|
||||||
|
func PermitInterface(luid uint64) (*Session, error) {
|
||||||
|
s, sublayerKey, err := newSession()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := addInterfaceFilter(s.engine, sublayerKey, fwpmLayerAleAuthRecvAcceptV4, luid); err != nil {
|
||||||
|
s.Close()
|
||||||
|
return nil, fmt.Errorf("add v4 filter: %w", err)
|
||||||
|
}
|
||||||
|
if err := addInterfaceFilter(s.engine, sublayerKey, fwpmLayerAleAuthRecvAcceptV6, luid); err != nil {
|
||||||
|
s.Close()
|
||||||
|
return nil, fmt.Errorf("add v6 filter: %w", err)
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// PermitUDPPort installs PERMIT filters at FWPM_LAYER_ALE_AUTH_RECV_ACCEPT_V4 and _V6 scoped to UDP traffic with the
|
||||||
|
// given local port. Inbound UDP to that port on any interface bypasses Windows Defender Firewall.
|
||||||
|
func PermitUDPPort(port uint16) (*Session, error) {
|
||||||
|
s, sublayerKey, err := newSession()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := addUDPPortFilter(s.engine, sublayerKey, fwpmLayerAleAuthRecvAcceptV4, port); err != nil {
|
||||||
|
s.Close()
|
||||||
|
return nil, fmt.Errorf("add v4 filter: %w", err)
|
||||||
|
}
|
||||||
|
if err := addUDPPortFilter(s.engine, sublayerKey, fwpmLayerAleAuthRecvAcceptV6, port); err != nil {
|
||||||
|
s.Close()
|
||||||
|
return nil, fmt.Errorf("add v6 filter: %w", err)
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSession() (*Session, windows.GUID, error) {
|
||||||
|
engine, err := openDynamicEngine()
|
||||||
|
if err != nil {
|
||||||
|
return nil, windows.GUID{}, err
|
||||||
|
}
|
||||||
|
sublayerKey, err := registerSublayer(engine)
|
||||||
|
if err != nil {
|
||||||
|
procFwpmEngineClose0.Call(engine)
|
||||||
|
return nil, windows.GUID{}, err
|
||||||
|
}
|
||||||
|
return &Session{engine: engine}, sublayerKey, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func openDynamicEngine() (uintptr, error) {
|
||||||
|
session := fwpmSession0{flags: fwpmSessionFlagDynamic}
|
||||||
|
var engine uintptr
|
||||||
|
r1, _, _ := procFwpmEngineOpen0.Call(
|
||||||
|
0, // serverName == NULL (local)
|
||||||
|
uintptr(rpcCAuthnWinNT),
|
||||||
|
0, // authIdentity == NULL
|
||||||
|
uintptr(unsafe.Pointer(&session)),
|
||||||
|
uintptr(unsafe.Pointer(&engine)),
|
||||||
|
)
|
||||||
|
if r1 != 0 {
|
||||||
|
return 0, fmt.Errorf("FwpmEngineOpen0: 0x%x", r1)
|
||||||
|
}
|
||||||
|
return engine, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// registerSublayer adds a session-scoped sublayer with a freshly generated GUID, weight 0xFFFF so its filters arbitrate
|
||||||
|
// above WDF's default sublayer. The sublayer is dynamic (no PERSISTENT flag) and goes away when the engine handle closes.
|
||||||
|
func registerSublayer(engine uintptr) (windows.GUID, error) {
|
||||||
|
key, err := windows.GenerateGUID()
|
||||||
|
if err != nil {
|
||||||
|
return windows.GUID{}, fmt.Errorf("GenerateGUID for sublayer: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
name, _ := windows.UTF16PtrFromString("Nebula WDF bypass sublayer")
|
||||||
|
desc, _ := windows.UTF16PtrFromString("Permit filters bypassing Windows Defender Firewall")
|
||||||
|
sl := fwpmSublayer0{
|
||||||
|
subLayerKey: key,
|
||||||
|
displayData: fwpmDisplayData0{name: name, description: desc},
|
||||||
|
weight: 0xFFFF,
|
||||||
|
}
|
||||||
|
r1, _, _ := procFwpmSubLayerAdd0.Call(
|
||||||
|
engine,
|
||||||
|
uintptr(unsafe.Pointer(&sl)),
|
||||||
|
0, // sd == NULL
|
||||||
|
)
|
||||||
|
if r1 != 0 {
|
||||||
|
return windows.GUID{}, fmt.Errorf("FwpmSubLayerAdd0: 0x%x", r1)
|
||||||
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func addInterfaceFilter(engine uintptr, sublayerKey, layer windows.GUID, luid uint64) error {
|
||||||
|
name, _ := windows.UTF16PtrFromString("Nebula allow interface inbound")
|
||||||
|
desc, _ := windows.UTF16PtrFromString("Permits inbound traffic on a nebula interface")
|
||||||
|
|
||||||
|
// luid must remain addressable through the syscall -- FWP_UINT64 is stored
|
||||||
|
// by pointer in the FWP_VALUE0 union.
|
||||||
|
cond := fwpmFilterCondition0{
|
||||||
|
fieldKey: fwpmConditionIPLocalInterface,
|
||||||
|
matchType: fwpMatchEqual,
|
||||||
|
conditionValue: fwpValue0{
|
||||||
|
type_: fwpUint64,
|
||||||
|
value: uintptr(unsafe.Pointer(&luid)),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
filter := fwpmFilter0{
|
||||||
|
// filterKey left zero: WFP assigns one when the filter is added.
|
||||||
|
displayData: fwpmDisplayData0{name: name, description: desc},
|
||||||
|
flags: fwpmFilterFlagClearActionRight,
|
||||||
|
layerKey: layer,
|
||||||
|
subLayerKey: sublayerKey,
|
||||||
|
weight: fwpValue0{type_: fwpUint8, value: uintptr(15)},
|
||||||
|
numFilterConditions: 1,
|
||||||
|
filterCondition: &cond,
|
||||||
|
action: fwpmAction0{actionType: fwpActionPermit},
|
||||||
|
}
|
||||||
|
|
||||||
|
r1, _, _ := procFwpmFilterAdd0.Call(
|
||||||
|
engine,
|
||||||
|
uintptr(unsafe.Pointer(&filter)),
|
||||||
|
0, // sd == NULL
|
||||||
|
0, // id == NULL
|
||||||
|
)
|
||||||
|
if r1 != 0 {
|
||||||
|
return fmt.Errorf("FwpmFilterAdd0: 0x%x", r1)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// addUDPPortFilter installs a PERMIT filter that matches (IP_PROTOCOL == UDP) AND (IP_LOCAL_PORT == port).
|
||||||
|
// FWP_UINT8 and FWP_UINT16 are <= 32 bits so they live inline in the FWP_VALUE0 union.
|
||||||
|
func addUDPPortFilter(engine uintptr, sublayerKey, layer windows.GUID, port uint16) error {
|
||||||
|
name, _ := windows.UTF16PtrFromString("Nebula allow UDP port inbound")
|
||||||
|
desc, _ := windows.UTF16PtrFromString("Permits inbound UDP to a nebula listener port")
|
||||||
|
|
||||||
|
conds := [2]fwpmFilterCondition0{
|
||||||
|
{
|
||||||
|
fieldKey: fwpmConditionIPProtocol,
|
||||||
|
matchType: fwpMatchEqual,
|
||||||
|
conditionValue: fwpValue0{
|
||||||
|
type_: fwpUint8,
|
||||||
|
value: uintptr(ipprotoUDP),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
fieldKey: fwpmConditionIPLocalPort,
|
||||||
|
matchType: fwpMatchEqual,
|
||||||
|
conditionValue: fwpValue0{
|
||||||
|
type_: fwpUint16,
|
||||||
|
value: uintptr(port),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
filter := fwpmFilter0{
|
||||||
|
displayData: fwpmDisplayData0{name: name, description: desc},
|
||||||
|
flags: fwpmFilterFlagClearActionRight,
|
||||||
|
layerKey: layer,
|
||||||
|
subLayerKey: sublayerKey,
|
||||||
|
weight: fwpValue0{type_: fwpUint8, value: uintptr(15)},
|
||||||
|
numFilterConditions: 2,
|
||||||
|
filterCondition: &conds[0],
|
||||||
|
action: fwpmAction0{actionType: fwpActionPermit},
|
||||||
|
}
|
||||||
|
|
||||||
|
r1, _, _ := procFwpmFilterAdd0.Call(
|
||||||
|
engine,
|
||||||
|
uintptr(unsafe.Pointer(&filter)),
|
||||||
|
0,
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
if r1 != 0 {
|
||||||
|
return fmt.Errorf("FwpmFilterAdd0: 0x%x", r1)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
package wire
|
||||||
|
|
||||||
|
// TunPacket is the unit a read from a tun device returns.
|
||||||
|
// On supported platforms, it may be a superpacket, but a single TunPacket will never have more than one destination.
|
||||||
|
type TunPacket struct {
|
||||||
|
// Bytes contains the actual packet
|
||||||
|
Bytes []byte
|
||||||
|
// Meta contains other information to help process the packet correctly, such as offsets for segmentation offloads
|
||||||
|
// Fields in Meta should be as portable/platform-agnostic as possible.
|
||||||
|
Meta struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
// PerSegment invokes fn once per segment of pkt.
|
||||||
|
// This is a stub implementation that does not actually support segmentation
|
||||||
|
func (t *TunPacket) PerSegment(fn func(seg []byte) error) error {
|
||||||
|
return fn(t.Bytes)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user